Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c5e0cbba6a |
@@ -220,20 +220,6 @@ DOUBAO_MAX_RETRIES=2
|
||||
DOUBAO_VISION_MODEL=doubao-1-5-vision-pro-250328
|
||||
DOUBAO_VISION_LITE_MODEL=doubao-1-5-vision-lite-250315
|
||||
DOUBAO_VISION_USE_LITE=true
|
||||
# Embedding 向量化模型(原 large-text-240915 已下线,用多模态 embedding)
|
||||
DOUBAO_EMBEDDING_MODEL=doubao-embedding-vision-251215
|
||||
|
||||
# ==================== 即梦(Jimeng)视觉 API —— 真人参考图兜底通道 (#2169) ====
|
||||
# 方舟 Seedance 走 B 端审核,真人参考图会被 50411 拦截;即梦走 C 端审核,普通真人照片可过审。
|
||||
# 需要在火山控制台开通即梦 cvtob 服务,使用 AK/SK(Region=cn-north-1, Service=cv)
|
||||
# 留空则真人拦截后直接返回错误提示,不会走即梦兜底。
|
||||
JIMENG_AK=
|
||||
JIMENG_SK=
|
||||
JIMENG_BASE_URL=https://visual.volcengineapi.com
|
||||
# 视频3.0 720P 首帧图生视频(1张图,稳定版;1080P标注下线中)
|
||||
JIMENG_REQ_KEY=jimeng_i2v_first_v30
|
||||
JIMENG_VIDEO_TIMEOUT=600
|
||||
JIMENG_VIDEO_POLL_INTERVAL=5
|
||||
|
||||
# ==================== 积分/会员系统 (#1895) ====================
|
||||
# 积分系统总开关:默认 false(暂停积分系统)。
|
||||
|
||||
@@ -1,87 +0,0 @@
|
||||
"""viral_video 动态积分定价 + 积分字段从 Integer 改为 Float (#2151)
|
||||
|
||||
Revision ID: 093
|
||||
Revises: 092_viral_video_heartbeat
|
||||
Create Date: 2026-10-02
|
||||
"""
|
||||
|
||||
import sqlalchemy as sa
|
||||
|
||||
from alembic import op
|
||||
|
||||
revision = "093"
|
||||
down_revision = "092_viral_video_heartbeat"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
inspector = sa.inspect(conn)
|
||||
|
||||
# 1) points_accounts 三列 Integer -> Float
|
||||
pa_cols = {c["name"]: c for c in inspector.get_columns("points_accounts")}
|
||||
for col in ("balance", "total_earned", "total_spent"):
|
||||
if col in pa_cols:
|
||||
op.alter_column(
|
||||
"points_accounts",
|
||||
col,
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 2) points_transactions amount/balance_after Integer -> Float
|
||||
pt_cols = {c["name"]: c for c in inspector.get_columns("points_transactions")}
|
||||
for col in ("amount", "balance_after"):
|
||||
if col in pt_cols:
|
||||
op.alter_column(
|
||||
"points_transactions",
|
||||
col,
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 3) users.points_balance Integer -> Float
|
||||
user_cols = {c["name"]: c for c in inspector.get_columns("users")}
|
||||
if "points_balance" in user_cols:
|
||||
op.alter_column(
|
||||
"users",
|
||||
"points_balance",
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 4) viral_video_jobs.credits_cost Integer -> Float
|
||||
vv_cols = {c["name"]: c for c in inspector.get_columns("viral_video_jobs")}
|
||||
if "credits_cost" in vv_cols:
|
||||
op.alter_column(
|
||||
"viral_video_jobs",
|
||||
"credits_cost",
|
||||
existing_type=sa.Integer(),
|
||||
type_=sa.Float(),
|
||||
existing_nullable=False,
|
||||
)
|
||||
|
||||
# 5) viral_video_jobs 新增列
|
||||
if "video_resolution" not in vv_cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("video_resolution", sa.String(20), nullable=False, server_default="720p"),
|
||||
)
|
||||
if "credits_prepaid" not in vv_cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("credits_prepaid", sa.Float(), nullable=False, server_default="0"),
|
||||
)
|
||||
if "credits_transaction_id" not in vv_cols:
|
||||
op.add_column(
|
||||
"viral_video_jobs",
|
||||
sa.Column("credits_transaction_id", sa.String(36), nullable=False, server_default=""),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
pass
|
||||
@@ -29,6 +29,8 @@ from app.services.ai_avatar_render_service import (
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -42,6 +44,7 @@ def _get_service(db: Session = Depends(get_db_session)) -> AiAvatarRenderService
|
||||
|
||||
|
||||
@router.post("", response_model=AiAvatarRenderJobResponse, status_code=201)
|
||||
@points_gate("ai_digital_human", per_unit=15)
|
||||
def create_render_job(
|
||||
body: CreateAiAvatarRenderRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
|
||||
@@ -27,6 +27,7 @@ from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
)
|
||||
from packages.application import ListGeneratedVideosByTaskUseCase
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
from packages.middleware.points_gate import points_gate
|
||||
from packages.shared.storage import get_shared_storage_service
|
||||
|
||||
from .templates_editor.dependencies import get_draft_plan_id, get_editor_services
|
||||
@@ -345,6 +346,7 @@ def _is_trusted_media_url(url: str) -> bool:
|
||||
|
||||
|
||||
@router.post("/generate-cover", response_model=GenerateCoverResponse)
|
||||
@points_gate("ai_cover")
|
||||
def generate_cover(
|
||||
body: GenerateCoverRequest,
|
||||
template_id: str = Query(..., description="模板 ID"),
|
||||
|
||||
@@ -41,6 +41,7 @@ from packages.application import (
|
||||
GetGenerationTaskUseCase,
|
||||
ListGeneratedVideosByTaskUseCase,
|
||||
)
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -269,6 +270,7 @@ def _variant_value(values: list[str], index: int, fallback: str = "") -> str:
|
||||
|
||||
|
||||
@router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, status_code=201)
|
||||
@points_gate("ai_video", quantity_field="preview_count")
|
||||
def create_preview_generation_task(
|
||||
request: CreatePreviewGenerationTaskRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
|
||||
@@ -163,6 +163,7 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None:
|
||||
return matched or None
|
||||
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -464,6 +465,7 @@ def _resolve_project_and_library(
|
||||
|
||||
|
||||
@router.post("/tasks", response_model=BatchGenerationTaskResponse)
|
||||
@points_gate("ai_video", quantity_field="count")
|
||||
def create_generation_task(
|
||||
request: CreateGenerationTaskRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
|
||||
@@ -12,9 +12,11 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import math
|
||||
from datetime import UTC
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.config import settings
|
||||
from app.dependencies import (
|
||||
get_db_session,
|
||||
get_voice_clone_profile_repository,
|
||||
@@ -30,6 +32,9 @@ from app.services.mediakit_client import MediaKitError
|
||||
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter()
|
||||
@@ -56,6 +61,37 @@ def create_lipsync_job(
|
||||
db: Session = Depends(get_db_session),
|
||||
svc: LipsyncService = Depends(_get_service),
|
||||
):
|
||||
user_id = current_user.user.id
|
||||
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_digital_human"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
# 口型同步:TTS 模式按 script_text 估时长(240字/分钟);音频直传按 audio_duration(秒→分钟)
|
||||
if body.audio_url and body.audio_duration and body.audio_duration > 0:
|
||||
est_minutes = max(1.0, math.ceil(body.audio_duration / 60.0))
|
||||
elif body.script_text:
|
||||
est_minutes = max(1.0, math.ceil(len(body.script_text) / 240))
|
||||
else:
|
||||
est_minutes = 1.0
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(current_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(current_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
"""提交对口型任务.
|
||||
|
||||
三种模式:
|
||||
@@ -65,8 +101,6 @@ def create_lipsync_job(
|
||||
- 预合成音频(#1845 新主路径):传 {video_url, audio_url, audio_duration, sentence_timings},
|
||||
后端同步ffprobe+写入timings+直接提交MediaKit(~2-3s)。
|
||||
"""
|
||||
user_id = current_user.user.id
|
||||
|
||||
try:
|
||||
job = svc.create_job(
|
||||
user_id=user_id,
|
||||
@@ -84,8 +118,18 @@ def create_lipsync_job(
|
||||
project_id=body.project_id,
|
||||
)
|
||||
except ValueError as exc:
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型 ValueError 退积分异常: err={refund_err}")
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
except MediaKitError as exc:
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型 MediaKitError 退积分异常: err={refund_err}")
|
||||
status_code = 502
|
||||
if exc.code in ("VoiceForbidden",):
|
||||
status_code = 403
|
||||
@@ -101,11 +145,24 @@ def create_lipsync_job(
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
logger.error("创建对口型任务异常: %s", exc, exc_info=True)
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型异常退积分异常: err={refund_err}")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"创建对口型任务失败: {exc}",
|
||||
) from exc
|
||||
|
||||
# 创建成功但状态为 failed(同步路径失败已抛异常到上面 except;此处处理 Celery 调度失败等)
|
||||
# 若任务已创建且状态为 failed,退费
|
||||
if _points_deducted > 0 and _points_svc is not None and getattr(job, "status", None) == "failed":
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"对口型任务失败退积分异常: job_id={job.id}, err={refund_err}")
|
||||
|
||||
return job
|
||||
|
||||
|
||||
@@ -119,14 +176,37 @@ def preview_tts(
|
||||
db: Session = Depends(get_db_session),
|
||||
svc: LipsyncService = Depends(_get_service),
|
||||
):
|
||||
user_id = current_user.user.id
|
||||
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_digital_human"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(body.script_text or "") / 240)) if body.script_text else 1.0
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(current_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(current_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
"""步骤1「生成配音」同步 TTS 预合成.
|
||||
|
||||
同步执行 TTS 合成 → 下载音频 → ffprobe 时长 → 句子时间戳计算,
|
||||
不创建 LipsyncJob、不转存 OSS,直接返回 CosyVoice 临时 URL(~24h 有效)。
|
||||
耗时约 2-3 秒。
|
||||
"""
|
||||
user_id = current_user.user.id
|
||||
|
||||
try:
|
||||
result = svc.preview_tts(
|
||||
user_id=user_id,
|
||||
@@ -138,6 +218,11 @@ def preview_tts(
|
||||
emotion=body.emotion,
|
||||
)
|
||||
except MediaKitError as exc:
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 预合成 MediaKitError 退积分异常: err={refund_err}")
|
||||
status_code = 400
|
||||
if exc.code in ("VoiceForbidden",):
|
||||
status_code = 403
|
||||
@@ -152,6 +237,11 @@ def preview_tts(
|
||||
) from exc
|
||||
except Exception as exc:
|
||||
logger.error("TTS 预合成异常: %s", exc, exc_info=True)
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 预合成异常退积分异常: err={refund_err}")
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"TTS 合成失败: {exc}",
|
||||
|
||||
@@ -145,22 +145,19 @@ def get_rules(
|
||||
def get_packages(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""查询可购买的积分包列表(读管理后台 credit_packages 表真实数据)。
|
||||
|
||||
仅返回 is_active=true;后台改价/启停后最多 30 秒生效。
|
||||
"""
|
||||
from packages.application.catalog.admin_catalog import get_points_packages
|
||||
|
||||
packages = [
|
||||
PointsPackageItem(
|
||||
code=row["code"],
|
||||
name=row["name"],
|
||||
points=row["points"],
|
||||
price_cents=row["price_cents"],
|
||||
unit_price=row["unit_price"],
|
||||
"""查询可购买的积分包列表。"""
|
||||
packages = []
|
||||
for code, pkg in POINTS_PACKAGES.items():
|
||||
unit_price = f"¥{pkg['price_cents'] / 100 / pkg['points']:.3f}/积分"
|
||||
packages.append(
|
||||
PointsPackageItem(
|
||||
code=code,
|
||||
name=pkg["name"],
|
||||
points=pkg["points"],
|
||||
price_cents=pkg["price_cents"],
|
||||
unit_price=unit_price,
|
||||
)
|
||||
)
|
||||
for row in get_points_packages()
|
||||
]
|
||||
mt = _member_type(current_user)
|
||||
discount = MEMBER_DISCOUNT.get(mt) if mt else None
|
||||
return PointsPackagesResponse(packages=packages, user_discount=discount)
|
||||
@@ -172,7 +169,17 @@ def check_points(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
):
|
||||
"""消费前检查余额是否足够。已下线/未知场景返回 cost=0(免费)。"""
|
||||
"""消费前检查余额是否足够。未知 scene_key 返回 400(而非 500)。"""
|
||||
if body.scene_key not in POINTS_SCENES:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail={
|
||||
"code": "UNKNOWN_SCENE",
|
||||
"message": f"未知场景: {body.scene_key}",
|
||||
"valid_scenes": sorted(POINTS_SCENES.keys()),
|
||||
},
|
||||
)
|
||||
|
||||
# 积分系统暂停(ENABLE_CREDIT_SYSTEM=false):所有场景直接放行,需 0 积分
|
||||
if not _credits_enabled():
|
||||
svc = _get_service()
|
||||
@@ -188,6 +195,13 @@ def check_points(
|
||||
is_mem = _is_member(current_user)
|
||||
mt = _member_type(current_user)
|
||||
|
||||
# 混剪场景先检查免费额度
|
||||
is_free_quota = False
|
||||
if body.scene_key == "ai_video" and not is_mem:
|
||||
svc = _get_service()
|
||||
if svc.check_daily_free_clip(current_user.user.id, db):
|
||||
is_free_quota = True
|
||||
|
||||
required = calculate_points_cost(
|
||||
body.scene_key,
|
||||
is_mem,
|
||||
@@ -201,11 +215,11 @@ def check_points(
|
||||
balance = account["balance"]
|
||||
|
||||
return PointsCheckResponse(
|
||||
allowed=balance >= required,
|
||||
allowed=is_free_quota or balance >= required,
|
||||
required_points=required,
|
||||
current_balance=balance,
|
||||
remaining_after=balance - required,
|
||||
is_free_quota=False,
|
||||
is_free_quota=is_free_quota,
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -44,6 +44,7 @@ from app.services.script_asr_service import (
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.middleware.points_gate import points_gate
|
||||
from packages.shared.ai_client import get_doubao_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -372,6 +373,7 @@ def douyin_diag():
|
||||
|
||||
|
||||
@router.post("/extract-from-douyin", response_model=ExtractFromDouyinResponse)
|
||||
@points_gate("douyin_extract")
|
||||
def extract_from_douyin(
|
||||
request: ExtractFromDouyinRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -495,6 +497,7 @@ def extract_from_douyin(
|
||||
|
||||
|
||||
@router.post("/ai-rewrite", response_model=AiRewriteResponse)
|
||||
@points_gate("ai_rewrite")
|
||||
def ai_rewrite(
|
||||
request: AiRewriteRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
@@ -534,6 +537,7 @@ def ai_rewrite(
|
||||
|
||||
|
||||
@router.post("/ai-generate-titles", response_model=AiGenerateTitlesResponse)
|
||||
@points_gate("ai_title")
|
||||
def ai_generate_titles(
|
||||
request: AiGenerateTitlesRequest,
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
|
||||
@@ -86,13 +86,33 @@ async def get_current_subscription(
|
||||
def list_membership_plans(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> dict[str, list[dict[str, Any]]]:
|
||||
"""查询可购买的会员套餐(读管理后台 plans 表真实数据)。
|
||||
"""查询所有会员档位(供前端会员购买页展示)。
|
||||
|
||||
仅返回 is_enabled=true 的套餐;后台启停/改价后最多 30 秒生效。
|
||||
返回 points 积分体系下的会员档位(月卡/季卡/年卡),含价格、时长、积分折扣等信息。
|
||||
"""
|
||||
from packages.application.catalog.admin_catalog import get_membership_plans
|
||||
from packages.domain.points_rules import MEMBER_DISCOUNT, MEMBERSHIP_PRICES
|
||||
|
||||
return {"plans": get_membership_plans()}
|
||||
plans: list[dict[str, Any]] = []
|
||||
for plan_id, info in MEMBERSHIP_PRICES.items():
|
||||
days = info["duration_days"]
|
||||
monthly_cents = round(info["price_cents"] * 30 / days)
|
||||
features: dict[str, Any] = {"max_resolution": "1080p"}
|
||||
if plan_id == MembershipType.MONTHLY:
|
||||
features.update({"free_clips_daily": 2})
|
||||
elif plan_id == MembershipType.QUARTERLY:
|
||||
features.update({"free_clips_daily": 5})
|
||||
elif plan_id == MembershipType.YEARLY:
|
||||
features.update({"free_clips_daily": "unlimited"})
|
||||
plans.append({
|
||||
"plan_id": plan_id,
|
||||
"name": info["name"],
|
||||
"price_cents": info["price_cents"],
|
||||
"monthly_price_cents": monthly_cents,
|
||||
"duration_days": days,
|
||||
"points_discount": MEMBER_DISCOUNT.get(plan_id, 1.0),
|
||||
"features": features,
|
||||
})
|
||||
return {"plans": plans}
|
||||
|
||||
|
||||
@router.get("/billing-records", response_model=list[BillingRecord])
|
||||
|
||||
@@ -4,12 +4,14 @@ from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import subprocess
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Any, Optional
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.config import settings
|
||||
from app.core.celery_app import celery_app
|
||||
from app.core.storage import get_storage_service
|
||||
from app.dependencies import (
|
||||
@@ -51,6 +53,8 @@ from packages.application.tts_job.use_cases import (
|
||||
)
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
from packages.domain.voice_presets import list_voices
|
||||
from packages.ports.asset_library_repository import AssetLibraryRepository
|
||||
from packages.ports.asset_repository import AssetRepository
|
||||
@@ -140,6 +144,31 @@ def synthesize(
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_voice"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
# 中文按 ~240 字/分钟粗估时长,至少按 1 分钟扣 1 分
|
||||
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(authenticated_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(authenticated_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
|
||||
# 解析 voice_id:前端可能传克隆音色 profile UUID(而非 CosyVoice voice_id),
|
||||
# 与 /tts/preview 保持一致:命中 profile → 校验归属 → 取 CosyVoice voice_id
|
||||
actual_voice_id = request.voice_id
|
||||
@@ -202,6 +231,7 @@ def synthesize(
|
||||
cosyvoice_service=cosyvoice_service,
|
||||
)
|
||||
|
||||
synthesis_error: Exception | None = None
|
||||
try:
|
||||
job = workflow.start_synthesis(job.id)
|
||||
except Exception as e:
|
||||
@@ -209,10 +239,18 @@ def synthesize(
|
||||
# 但 DB 异常、网络异常等意外错误可能逃逸。
|
||||
# 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。
|
||||
logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True)
|
||||
synthesis_error = e
|
||||
try:
|
||||
job = workflow.process_synthesis_failure(job.id, str(e))
|
||||
except Exception as inner_e:
|
||||
logger.error(f"标记 TTS job 失败时出错: job_id={job.id}, error={inner_e}")
|
||||
# 合成失败且已扣积分 → 退费
|
||||
if synthesis_error is not None and _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 合失败退积分异常: job_id={job.id}, err={refund_err}")
|
||||
|
||||
# 若任务处于 processing 状态(异步模式),触发 Celery 后台轮询
|
||||
if job.status.value == "processing":
|
||||
# 分段合成任务 vs 普通单段任务
|
||||
@@ -231,6 +269,13 @@ def synthesize(
|
||||
workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}")
|
||||
except Exception as inner_e:
|
||||
logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={inner_e}")
|
||||
# 调度失败退费
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"Celery 调度失败退积分异常: job_id={job.id}, err={refund_err}")
|
||||
|
||||
return TTSSynthesizeResponse(
|
||||
job_id=job.id,
|
||||
status=job.status,
|
||||
@@ -565,6 +610,31 @@ def preview_tts(
|
||||
用于前端预览配音效果,限制文本长度 200 字以内。
|
||||
支持预设音色和克隆音色:克隆音色传的是 profile UUID,需解析为 CosyVoice voice_id。
|
||||
"""
|
||||
user_id = authenticated_user.user.id
|
||||
# ── 积分扣点(#1895 P2) ──
|
||||
_points_deducted = 0
|
||||
_points_scene = "ai_voice"
|
||||
_points_svc = PointsService() if settings.points_enabled else None
|
||||
if _points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(request.text) / 240))
|
||||
_points_deducted = calculate_points_cost(
|
||||
_points_scene,
|
||||
is_member=getattr(authenticated_user.user, "is_member", False),
|
||||
duration_minutes=est_minutes,
|
||||
member_type=getattr(authenticated_user.user, "member_type", None),
|
||||
)
|
||||
_deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db)
|
||||
if not _deduct_res["success"]:
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}",
|
||||
"required": _points_deducted,
|
||||
"balance": _deduct_res["balance"],
|
||||
},
|
||||
)
|
||||
|
||||
# 解析 voice_id:前端可能传 VoiceCloneProfile UUID 或预设音色 ID
|
||||
actual_voice_id = request.voice_id
|
||||
profile = voice_clone_repo.get(request.voice_id)
|
||||
@@ -594,6 +664,12 @@ def preview_tts(
|
||||
language=getattr(request, "language", "zh-CN"),
|
||||
)
|
||||
except (CosyVoiceError, ValueError) as e:
|
||||
# 合成失败退费
|
||||
if _points_deducted > 0 and _points_svc is not None:
|
||||
try:
|
||||
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||
except Exception as refund_err:
|
||||
logger.warning(f"TTS 预览失败退积分异常: {refund_err}")
|
||||
if isinstance(e, CosyVoiceError):
|
||||
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=f"TTS 合成失败: {e}") from e
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
|
||||
|
||||
@@ -32,11 +32,7 @@ from app.schemas.viral_video import (
|
||||
ConfirmCopyRequest,
|
||||
ConfirmIntentRequest,
|
||||
CreateViralVideoRequest,
|
||||
CreditsFormulaBreakdown,
|
||||
EstimateCreditsRequest,
|
||||
EstimateCreditsResponse,
|
||||
GenerateCopyRequest,
|
||||
RetryViralVideoRequest,
|
||||
StyleTemplateListResponse,
|
||||
StyleTemplateResponse,
|
||||
ViralVideoHistoryResponse,
|
||||
@@ -49,9 +45,7 @@ from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
SQLAlchemyViralVideoStyleTemplateRepository,
|
||||
)
|
||||
from packages.domain.points_rules import list_viral_video_models
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -146,9 +140,7 @@ def _to_response(job) -> ViralVideoJobResponse:
|
||||
video_model=getattr(job, "video_model", "") or "",
|
||||
intent_result=job.intent_result,
|
||||
result_video_url=job.result_video_url,
|
||||
video_resolution=getattr(job, "video_resolution", "720p") or "720p",
|
||||
credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0),
|
||||
credits_cost=float(getattr(job, "credits_cost", 0) or 0),
|
||||
credits_cost=job.credits_cost,
|
||||
error_msg=job.error_msg,
|
||||
retry_count=job.retry_count,
|
||||
started_at=job.started_at,
|
||||
@@ -201,7 +193,6 @@ def create_viral_video(
|
||||
voice_source=getattr(request, "voice_source", "") or "",
|
||||
video_ratio=getattr(request, "video_ratio", "9:16") or "9:16",
|
||||
video_model=getattr(request, "video_model", "") or "",
|
||||
video_resolution=getattr(request, "video_resolution", "720p") or "720p",
|
||||
copy_result=None,
|
||||
)
|
||||
|
||||
@@ -244,7 +235,6 @@ def analyze_images(
|
||||
voice_source=request.voice_source or "",
|
||||
video_ratio=request.video_ratio or "9:16",
|
||||
video_model=request.video_model or "",
|
||||
video_resolution=getattr(request, "video_resolution", "720p") or "720p",
|
||||
duration=request.duration or 15,
|
||||
)
|
||||
repo.save(job)
|
||||
@@ -308,7 +298,6 @@ def generate_copy(
|
||||
job.voice_source = request.voice_source or job.voice_source
|
||||
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"
|
||||
|
||||
job.resume_from_image_analyzed()
|
||||
repo.update(job)
|
||||
@@ -341,41 +330,6 @@ def confirm_copy(
|
||||
if job.status != ViralVideoStatus.COPY_GENERATED:
|
||||
raise HTTPException(status_code=409, detail=f"任务当前状态 {job.status} 不能确认文案(需 copy_generated)")
|
||||
|
||||
# 积分预扣(已扣过/重试任务跳过)
|
||||
from app.config import settings as _settings
|
||||
|
||||
if _settings.points_enabled:
|
||||
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:
|
||||
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
w, h = resolve_video_dimensions(
|
||||
getattr(job, "video_resolution", "720p") or "720p",
|
||||
job.video_ratio or "9:16",
|
||||
)
|
||||
est_credits = calculate_viral_video_credits(
|
||||
int(job.duration or 15), w, h, job.video_model or "seedance-2.5"
|
||||
)
|
||||
svc = PointsService()
|
||||
res = svc.deduct_viral_video(authenticated_user.user.id, est_credits, job.id, session)
|
||||
if not res.get("success"):
|
||||
balance = res.get("balance", 0)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"积分不足,需要 {est_credits} 积分,当前余额 {balance}",
|
||||
"required": est_credits,
|
||||
"balance": balance,
|
||||
},
|
||||
)
|
||||
job.credits_prepaid = est_credits
|
||||
job.credits_transaction_id = res.get("transaction_id", "") or ""
|
||||
repo.update(job)
|
||||
|
||||
job.resume_from_copy_generated(edited_copy=request.edited_copy or None)
|
||||
repo.update(job)
|
||||
|
||||
@@ -390,38 +344,6 @@ def confirm_copy(
|
||||
return _to_response(job)
|
||||
|
||||
|
||||
@router.post("/estimate-credits", response_model=EstimateCreditsResponse)
|
||||
def estimate_credits(
|
||||
request: EstimateCreditsRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EstimateCreditsResponse:
|
||||
"""爆款视频积分预估(纯计算,不扣费、不创建任务)。
|
||||
|
||||
返回 estimated_credits 与 formula_breakdown(tokens / video_cost / fixed_cost /
|
||||
profit_multiplier / model_price / width / height / fps),便于前端展示计费明细。
|
||||
同时兼容前端传 model 或 video_model、resolution 或 video_resolution、ratio 或 video_ratio。
|
||||
"""
|
||||
from packages.domain.points_rules import (
|
||||
calculate_viral_video_credits_with_breakdown,
|
||||
resolve_video_dimensions,
|
||||
)
|
||||
|
||||
model = (request.model or "").strip() or "seedance-2.5"
|
||||
resolution = (request.resolution or "").strip() or "720p"
|
||||
ratio = (request.ratio or "").strip() or "9:16"
|
||||
duration = int(request.duration or 15)
|
||||
|
||||
w, h = resolve_video_dimensions(resolution, ratio)
|
||||
credits, bd = calculate_viral_video_credits_with_breakdown(
|
||||
duration,
|
||||
w,
|
||||
h,
|
||||
model,
|
||||
)
|
||||
breakdown = CreditsFormulaBreakdown(**bd)
|
||||
return EstimateCreditsResponse(estimated_credits=credits, formula_breakdown=breakdown)
|
||||
|
||||
|
||||
@router.get("/history", response_model=ViralVideoHistoryResponse)
|
||||
def list_viral_video_history(
|
||||
limit: int = 50,
|
||||
@@ -456,17 +378,6 @@ def list_style_templates(
|
||||
return StyleTemplateListResponse(items=items)
|
||||
|
||||
|
||||
@router.get("/models")
|
||||
def list_available_models() -> dict:
|
||||
"""返回爆款视频可用模型列表(供前端模型选择器使用)。"""
|
||||
dashscope_available = get_dashscope_client() is not None
|
||||
models = list_viral_video_models(
|
||||
include_placeholder=False,
|
||||
dashscope_available=dashscope_available,
|
||||
)
|
||||
return {"models": models}
|
||||
|
||||
|
||||
@router.get("/{job_id}", response_model=ViralVideoJobResponse)
|
||||
def get_viral_video_job(
|
||||
job_id: str,
|
||||
@@ -486,17 +397,10 @@ def get_viral_video_job(
|
||||
@router.post("/{job_id}/retry", response_model=ViralVideoJobResponse)
|
||||
def retry_viral_video_job(
|
||||
job_id: str,
|
||||
request: RetryViralVideoRequest | None = None,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> ViralVideoJobResponse:
|
||||
"""重试失败的爆款视频任务(也支持对僵尸/超时 running 任务强制重置后重试)。
|
||||
|
||||
可选 body (RetryViralVideoRequest):若传入新的 duration/video_resolution/video_ratio/
|
||||
video_model,会重新预估积分并与原 credits_prepaid 做差额多退少补(不足抛 402 阻止重试);
|
||||
不传 body 或参数无变化时,保持原参数、原预扣金额不变,仅重置状态并入队。
|
||||
credits_prepaid 为 0 的老任务首次重试会走预扣流程(与 confirm-copy 一致)。
|
||||
"""
|
||||
"""重试失败的爆款视频任务(也支持对僵尸/超时 running 任务强制重置后重试)。"""
|
||||
from datetime import datetime, timezone
|
||||
|
||||
repo = _get_job_repo(session)
|
||||
@@ -517,104 +421,6 @@ def retry_viral_video_job(
|
||||
if job.status != ViralVideoStatus.FAILED and not is_stale_running:
|
||||
raise HTTPException(status_code=409, detail="只有失败或超时的任务可以重试")
|
||||
|
||||
# ── 参数变更检测 + 积分多退少补 ──────────────────────────────────────
|
||||
req = request or RetryViralVideoRequest()
|
||||
new_duration = req.duration
|
||||
new_resolution = (req.video_resolution or "").strip() or None
|
||||
new_ratio = (req.video_ratio or "").strip() or None
|
||||
new_model = (req.video_model or "").strip() or None
|
||||
|
||||
old_duration = int(getattr(job, "duration", 15) or 15)
|
||||
old_resolution = (getattr(job, "video_resolution", "720p") or "720p").strip() or "720p"
|
||||
old_ratio = (getattr(job, "video_ratio", "9:16") or "9:16").strip() or "9:16"
|
||||
old_model = (getattr(job, "video_model", "") or "").strip()
|
||||
|
||||
# 仅当有任意字段传入且值不同才算"参数变更"
|
||||
param_changed = bool(
|
||||
(new_duration is not None and int(new_duration) != old_duration)
|
||||
or (new_resolution is not None and new_resolution != old_resolution)
|
||||
or (new_ratio is not None and new_ratio != old_ratio)
|
||||
or (new_model is not None and new_model != old_model)
|
||||
)
|
||||
|
||||
from app.config import settings as _settings
|
||||
|
||||
need_points_settle = False
|
||||
new_est = 0.0
|
||||
if _settings.points_enabled and param_changed:
|
||||
from packages.domain.points_rules import (
|
||||
calculate_viral_video_credits_with_breakdown,
|
||||
resolve_video_dimensions,
|
||||
)
|
||||
|
||||
eff_dur = int(new_duration if new_duration is not None else old_duration)
|
||||
eff_res = new_resolution if new_resolution is not None else old_resolution
|
||||
eff_ratio = new_ratio if new_ratio is not None else old_ratio
|
||||
eff_model = new_model if new_model is not None else (old_model or "seedance-2.5")
|
||||
w, h = resolve_video_dimensions(eff_res, eff_ratio)
|
||||
new_est, _ = calculate_viral_video_credits_with_breakdown(eff_dur, w, h, eff_model or "seedance-2.5")
|
||||
need_points_settle = True
|
||||
|
||||
# 写入新参数(即使不开 points 也要允许用户重试时改参数)
|
||||
if new_duration is not None:
|
||||
job.duration = max(5, min(30, int(new_duration)))
|
||||
if new_resolution is not None:
|
||||
job.video_resolution = new_resolution
|
||||
if new_ratio is not None:
|
||||
job.video_ratio = new_ratio
|
||||
if new_model is not None:
|
||||
job.video_model = new_model
|
||||
|
||||
if need_points_settle:
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
old_prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
svc = PointsService()
|
||||
diff = round(new_est - old_prepaid, 2)
|
||||
if abs(diff) >= 0.01:
|
||||
if diff > 0:
|
||||
# 新预扣更多:补扣差额
|
||||
res = svc.deduct_viral_video(authenticated_user.user.id, diff, job.id, session)
|
||||
if not res.get("success"):
|
||||
balance = res.get("balance", 0)
|
||||
raise HTTPException(
|
||||
status_code=402,
|
||||
detail={
|
||||
"code": "INSUFFICIENT_POINTS",
|
||||
"message": f"重试参数变更后需补扣 {diff} 积分,余额不足(当前 {balance},需 {new_est})",
|
||||
"required": new_est,
|
||||
"balance": balance,
|
||||
"delta": diff,
|
||||
},
|
||||
)
|
||||
job.credits_prepaid = round(old_prepaid + diff, 2)
|
||||
logger.info(
|
||||
"[爆款视频][retry] 补扣差额 job_id=%s diff=%.2f new_prepaid=%.2f",
|
||||
job.id,
|
||||
diff,
|
||||
job.credits_prepaid,
|
||||
)
|
||||
else:
|
||||
# 新预扣更少:退还差额
|
||||
refund = round(-diff, 2)
|
||||
txn_id = getattr(job, "credits_transaction_id", "") or ""
|
||||
svc.refund_points(
|
||||
user_id=authenticated_user.user.id,
|
||||
amount=refund,
|
||||
source="viral_video",
|
||||
db=session,
|
||||
ref_id=txn_id or job.id,
|
||||
description="爆款视频重试参数变更退费",
|
||||
)
|
||||
job.credits_prepaid = round(old_prepaid - refund, 2)
|
||||
logger.info(
|
||||
"[爆款视频][retry] 退还差额 job_id=%s refund=%.2f new_prepaid=%.2f",
|
||||
job.id,
|
||||
refund,
|
||||
job.credits_prepaid,
|
||||
)
|
||||
# 差额为 0 则不调整
|
||||
|
||||
# 重置状态
|
||||
job.retry_count += 1
|
||||
job.status = ViralVideoStatus.PENDING
|
||||
@@ -629,13 +435,7 @@ def retry_viral_video_job(
|
||||
# 重新入队
|
||||
try:
|
||||
celery_app.send_task("worker.run_viral_video_pipeline", args=[job.id])
|
||||
logger.info(
|
||||
"[爆款视频] 重试入队: job_id=%s retry_count=%d stale=%s params_changed=%s",
|
||||
job.id,
|
||||
job.retry_count,
|
||||
is_stale_running,
|
||||
param_changed,
|
||||
)
|
||||
logger.info("[爆款视频] 重试入队: job_id=%s retry_count=%d stale=%s", job.id, job.retry_count, is_stale_running)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True)
|
||||
job.mark_failed(f"重试入队失败: {e}")
|
||||
|
||||
@@ -13,9 +13,9 @@ from pydantic import BaseModel, Field
|
||||
class PointsBalanceResponse(BaseModel):
|
||||
"""积分余额 + 会员状态"""
|
||||
|
||||
balance: float = Field(..., description="当前积分余额")
|
||||
total_earned: float = Field(..., description="累计获得积分")
|
||||
total_spent: float = Field(..., description="累计消耗积分")
|
||||
balance: int = Field(..., description="当前积分余额")
|
||||
total_earned: int = Field(..., description="累计获得积分")
|
||||
total_spent: int = Field(..., description="累计消耗积分")
|
||||
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="会员到期时间")
|
||||
@@ -30,8 +30,8 @@ class PointsTransactionItem(BaseModel):
|
||||
id: str
|
||||
type: str = Field(..., description="类型: add/deduct")
|
||||
source: str = Field(..., description="来源场景")
|
||||
amount: float
|
||||
balance_after: float
|
||||
amount: int
|
||||
balance_after: int
|
||||
description: str = ""
|
||||
ref_id: str = ""
|
||||
created_at: Optional[str] = None
|
||||
@@ -99,9 +99,9 @@ class PointsCheckResponse(BaseModel):
|
||||
"""消费前余额检查响应"""
|
||||
|
||||
allowed: bool
|
||||
required_points: float
|
||||
current_balance: float
|
||||
remaining_after: float
|
||||
required_points: int
|
||||
current_balance: int
|
||||
remaining_after: int
|
||||
is_free_quota: bool = False
|
||||
|
||||
|
||||
@@ -112,7 +112,7 @@ class PointsDeductRequest(BaseModel):
|
||||
"""积分扣减请求"""
|
||||
|
||||
scene_key: str
|
||||
amount: float
|
||||
amount: int
|
||||
description: Optional[str] = ""
|
||||
ref_id: Optional[str] = ""
|
||||
|
||||
@@ -170,7 +170,7 @@ class MembershipStatusResponse(BaseModel):
|
||||
is_member: bool
|
||||
member_type: Optional[str] = None
|
||||
member_expires_at: Optional[datetime] = None
|
||||
points_balance: float
|
||||
points_balance: int
|
||||
max_resolution: str = Field(
|
||||
default="1080p",
|
||||
description="可用最高分辨率: 720p(free) / 1080p(paid)",
|
||||
|
||||
@@ -22,7 +22,6 @@ VALID_STAGES = (
|
||||
)
|
||||
VALID_VIDEO_RATIOS = ("9:16", "16:9", "1:1", "4:3", "3:4", "21:9")
|
||||
VALID_DURATIONS = (5, 10, 15, 20, 25, 30)
|
||||
VALID_VIDEO_RESOLUTIONS = ("480p", "720p", "1080p", "普清", "高清", "超清")
|
||||
|
||||
|
||||
# -- 编导脚本结构(v1.6) --
|
||||
@@ -89,7 +88,6 @@ class CreateViralVideoRequest(BaseModel):
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
|
||||
@field_validator("fusion_level")
|
||||
@classmethod
|
||||
@@ -119,7 +117,6 @@ class AnalyzeImagesRequest(BaseModel):
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
duration: int = Field(default=15, ge=5, le=30)
|
||||
|
||||
|
||||
@@ -144,7 +141,6 @@ class GenerateCopyRequest(BaseModel):
|
||||
voice_source: str = ""
|
||||
video_ratio: str = "9:16"
|
||||
video_model: str = ""
|
||||
video_resolution: str = "720p"
|
||||
|
||||
@field_validator("fusion_level")
|
||||
@classmethod
|
||||
@@ -222,9 +218,7 @@ class ViralVideoJobResponse(BaseModel):
|
||||
video_model: str = ""
|
||||
intent_result: dict | None = None
|
||||
result_video_url: str = ""
|
||||
video_resolution: str = "720p"
|
||||
credits_prepaid: float = 0.0
|
||||
credits_cost: float = 0.0
|
||||
credits_cost: int = 0
|
||||
error_msg: str = ""
|
||||
retry_count: int = 0
|
||||
started_at: datetime | None = None
|
||||
@@ -256,57 +250,6 @@ class AnalyzeStyleResponse(BaseModel):
|
||||
style_guide: dict | None = None
|
||||
|
||||
|
||||
# -- 积分预估 --
|
||||
|
||||
|
||||
class EstimateCreditsRequest(BaseModel):
|
||||
"""爆款视频积分预估请求。
|
||||
|
||||
前端可传 model 或 video_model(兼容老字段);resolution/ratio/duration 为预估所需参数。
|
||||
"""
|
||||
|
||||
model: str = Field(default="", alias="video_model")
|
||||
resolution: str = Field(default="720p", alias="video_resolution")
|
||||
ratio: str = Field(default="9:16", alias="video_ratio")
|
||||
duration: int = Field(default=15, ge=5, le=30)
|
||||
|
||||
model_config = {"populate_by_name": True}
|
||||
|
||||
|
||||
class CreditsFormulaBreakdown(BaseModel):
|
||||
"""爆款视频积分计费公式明细(前端展示用)。"""
|
||||
|
||||
tokens: float = Field(..., description="估算视频 tokens 数 (duration*width*height*fps/1024)")
|
||||
video_cost: float = Field(..., description="视频生成成本(元)= tokens/1e6 * model_price")
|
||||
fixed_cost: float = Field(..., description="固定成本(元),含 VLM/LLM/TTS/OSS/服务器")
|
||||
profit_multiplier: float = Field(..., description="利润系数(默认 1.3)")
|
||||
model_price: float = Field(..., description="模型单价(元/百万 tokens)")
|
||||
width: int = Field(..., description="视频宽度像素")
|
||||
height: int = Field(..., description="视频高度像素")
|
||||
fps: int = Field(..., description="视频帧率")
|
||||
|
||||
|
||||
class EstimateCreditsResponse(BaseModel):
|
||||
"""爆款视频积分预估响应。"""
|
||||
|
||||
estimated_credits: float
|
||||
formula_breakdown: CreditsFormulaBreakdown = Field(..., description="计费公式明细")
|
||||
|
||||
|
||||
class RetryViralVideoRequest(BaseModel):
|
||||
"""重试爆款视频任务的请求体(可选,允许改参数重新预估积分多退少补)。
|
||||
|
||||
不传 body 或字段全缺省:保持原参数、不重新扣点,走默认重置+入队逻辑。
|
||||
传入新的 duration/video_resolution/video_ratio/video_model:重新预估积分,
|
||||
与原 credits_prepaid 比较后多退少补(差额补扣不足抛 402)。
|
||||
"""
|
||||
|
||||
duration: int | None = Field(default=None, ge=5, le=30, description="重试时新的视频时长(秒)")
|
||||
video_resolution: str | None = Field(default=None, description="重试时新的分辨率,如 720p/1080p")
|
||||
video_ratio: str | None = Field(default=None, description="重试时新的画幅比,如 9:16/16:9")
|
||||
video_model: str | None = Field(default=None, description="重试时新的视频模型,如 seedance-2.5")
|
||||
|
||||
|
||||
# -- WebSocket 事件 Schema --
|
||||
|
||||
|
||||
|
||||
@@ -11,12 +11,14 @@
|
||||
存储路径与元信息约定),返回 asset_id —— 下游仍以 voice_library_id(实为
|
||||
audio asset id)消费,渲染链路零改动。
|
||||
|
||||
积分扣点与 /tts 合成端点保持一致(ai_voice 场景),失败退费。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import math
|
||||
import subprocess
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
@@ -30,10 +32,13 @@ from packages.application.cosyvoice_service import CosyVoiceService
|
||||
from packages.application.tts_job.use_cases import CreateTTSJobUseCase
|
||||
from packages.application.tts_job.workflow import TTSWorkflowService
|
||||
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
from packages.domain.points_service import PointsService
|
||||
from packages.shared.storage import SharedStorageService
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_POINTS_SCENE = "ai_voice"
|
||||
_SYNTH_TIMEOUT = 180.0 # 叙事配音在 HTTP 请求内同步等待,长文案分段合成时留出余量
|
||||
_CONTENT_TYPE_MAP = {"mp3": "audio/mpeg", "wav": "audio/wav", "pcm": "audio/pcm", "opus": "audio/opus"}
|
||||
|
||||
@@ -268,6 +273,24 @@ def prepare_narrative_voice(
|
||||
voice_clone_repository=voice_clone_repository,
|
||||
)
|
||||
|
||||
# 积分扣点(与 /tts 合成端点同口径),失败时在合成失败分支退费
|
||||
points_svc = PointsService() if points_enabled else None
|
||||
points_deducted = 0
|
||||
if points_svc is not None:
|
||||
est_minutes = max(1.0, math.ceil(len(content) / 240))
|
||||
points_deducted = calculate_points_cost(
|
||||
_POINTS_SCENE,
|
||||
is_member=is_member,
|
||||
duration_minutes=est_minutes,
|
||||
member_type=member_type,
|
||||
)
|
||||
deduct_res = points_svc.deduct_points(user_id, points_deducted, _POINTS_SCENE, db)
|
||||
if not deduct_res["success"]:
|
||||
raise NarrativeError(
|
||||
f"积分不足,需要 {points_deducted} 积分,当前余额 {deduct_res['balance']}",
|
||||
status_code=402,
|
||||
)
|
||||
|
||||
use_case = CreateTTSJobUseCase(tts_repository)
|
||||
job = use_case.execute(
|
||||
user_id=user_id,
|
||||
@@ -288,9 +311,19 @@ def prepare_narrative_voice(
|
||||
workflow.process_synthesis_failure(job.id, str(e))
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("标记叙事 TTS job 失败出错: job_id=%s", job.id, exc_info=True)
|
||||
if points_deducted and points_svc is not None:
|
||||
try:
|
||||
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("叙事 TTS 失败退积分异常: job_id=%s", job.id, exc_info=True)
|
||||
raise NarrativeError(f"配音合成失败:{e}", status_code=502) from e
|
||||
|
||||
if not job.is_completed:
|
||||
if points_deducted and points_svc is not None:
|
||||
try:
|
||||
points_svc.refund_points(user_id, points_deducted, _POINTS_SCENE, db, ref_id=job.id)
|
||||
except Exception: # noqa: BLE001
|
||||
logger.warning("叙事 TTS 未完成退积分异常: job_id=%s", job.id, exc_info=True)
|
||||
raise NarrativeError("配音合成未完成,请稍后重试", status_code=504)
|
||||
|
||||
asset = _save_tts_job_as_voice_asset(
|
||||
|
||||
@@ -9,8 +9,6 @@ import type {
|
||||
AnalyzeImagesRequest,
|
||||
GenerateCopyRequest,
|
||||
ConfirmCopyRequest,
|
||||
ViralVideoModel,
|
||||
ViralVideoModelsResponse,
|
||||
} from "./types"
|
||||
|
||||
/** 创建爆款视频任务 */
|
||||
@@ -53,29 +51,6 @@ export function analyzeViralStyle(id: string) {
|
||||
return apiClient.post<ViralVideoJob>(`/viral-video/${id}/analyze-style`).then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 动态预估积分消耗(STEP3 参数变化时调用) */
|
||||
export function estimateViralVideoCredits(params: {
|
||||
video_model: string
|
||||
resolution: string
|
||||
video_ratio: string
|
||||
duration: number
|
||||
}) {
|
||||
return apiClient
|
||||
.post<{ estimated_credits: number }>("/viral-video/estimate-credits", params)
|
||||
.then((r) => r.data)
|
||||
}
|
||||
|
||||
/** 获取支持的视频模型列表(GET /viral-video/models)。后端返回 {models: [...]} 包装 */
|
||||
export function getViralVideoModels() {
|
||||
return apiClient.get<ViralVideoModelsResponse>("/viral-video/models").then((r) => {
|
||||
const data = r.data as ViralVideoModelsResponse | ViralVideoModel[] | null | undefined
|
||||
if (Array.isArray(data)) return data
|
||||
if (data && Array.isArray((data as ViralVideoModelsResponse).models)) {
|
||||
return (data as ViralVideoModelsResponse).models
|
||||
}
|
||||
return []
|
||||
})
|
||||
}
|
||||
/** ── 三步拆分:前端 mock 辅助函数(后端新接口上线后可替换) ── */
|
||||
|
||||
/**
|
||||
|
||||
@@ -280,29 +280,10 @@ export interface GenerateCopyRequest {
|
||||
video_model?: string
|
||||
}
|
||||
|
||||
/** 视频模型描述(GET /viral-video/models) */
|
||||
export interface ViralVideoModel {
|
||||
key: string
|
||||
display_name: string
|
||||
supports_audio: boolean
|
||||
supported_resolutions: string[]
|
||||
max_duration: number
|
||||
/** 计费模式(可选):per_second / per_video / token 等 */
|
||||
billing_mode?: string
|
||||
is_default?: boolean
|
||||
}
|
||||
|
||||
/** GET /viral-video/models 响应包装 */
|
||||
export interface ViralVideoModelsResponse {
|
||||
models: ViralVideoModel[]
|
||||
}
|
||||
|
||||
/** v1.6 阶段3请求:用户确认/编辑口播文案后开始单次 Seedance 出片(POST /viral-video/{id}/confirm-copy) */
|
||||
export interface ConfirmCopyRequest {
|
||||
/** 用户编辑后的口播文案;为空则使用 AI 生成的 voiceover_script */
|
||||
edited_copy?: string
|
||||
/** 视频模型 key,覆盖默认 */
|
||||
video_model?: string
|
||||
}
|
||||
|
||||
/** 旧分镜片段结构(保留兼容;新代码请使用 ShotScript) */
|
||||
|
||||
@@ -80,8 +80,6 @@ const AiAvatarPage: React.FC = () => {
|
||||
const [finalizeLoading, setFinalizeLoading] = useState(false)
|
||||
|
||||
/* ── 对口型轮询 ── */
|
||||
/** 对口型轮询总时长上限(10分钟):超过后停止轮询并提示去历史记录查看 */
|
||||
const LIPSYNC_POLL_MAX_MS = 10 * 60 * 1000
|
||||
const lipsyncTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
/* ── 渲染进度轮询 ── */
|
||||
const renderTimerRef = useRef<ReturnType<typeof setInterval> | null>(null)
|
||||
@@ -272,24 +270,7 @@ const AiAvatarPage: React.FC = () => {
|
||||
// 如果是预合成模式,后端会同步把状态置为 submitted(甚至可能已返回 running),
|
||||
// 但仍需轮询等 completed
|
||||
if (lipsyncTimerRef.current) clearInterval(lipsyncTimerRef.current)
|
||||
// 轮询间隔 5 秒;单请求超时 5 分钟(见 api/aiAvatar.ts);总轮询上限 10 分钟
|
||||
// 单次请求失败/超时不中断轮询,继续下一轮;超过总上限后停止并提示用户去历史记录查看
|
||||
lipsyncTimerRef.current = setInterval(async () => {
|
||||
// 总时长保护:超过 10 分钟停止轮询
|
||||
if (Date.now() - lipsyncStartAtRef.current > LIPSYNC_POLL_MAX_MS) {
|
||||
if (lipsyncTimerRef.current) {
|
||||
clearInterval(lipsyncTimerRef.current)
|
||||
lipsyncTimerRef.current = null
|
||||
}
|
||||
if (lipsyncTickRef.current) {
|
||||
clearInterval(lipsyncTickRef.current)
|
||||
lipsyncTickRef.current = null
|
||||
}
|
||||
setLipsyncStatus("failed")
|
||||
setLipsyncErrorMessage("渲染时间较长,请稍后在历史记录中查看")
|
||||
message.warning("对口型渲染时间较长,已停止自动刷新,请稍后在历史记录中查看")
|
||||
return
|
||||
}
|
||||
try {
|
||||
const updated = await getLipsyncJob(job.id)
|
||||
state.setLipsyncJob(updated)
|
||||
@@ -315,10 +296,9 @@ const AiAvatarPage: React.FC = () => {
|
||||
setLipsyncErrorMessage(updated.error_message || "对口型生成失败")
|
||||
}
|
||||
} catch (err) {
|
||||
// 单次轮询失败(含 timeout):不中断轮询,打印日志后等下一轮
|
||||
console.warn("[对口型] 轮询请求失败,将继续下一轮:", err)
|
||||
console.error("[对口型] 轮询错误:", err)
|
||||
}
|
||||
}, 5000)
|
||||
}, 3000)
|
||||
} catch (err) {
|
||||
console.error("[对口型] 创建失败:", {
|
||||
status: (err as { response?: { status?: number } })?.response?.status,
|
||||
|
||||
@@ -72,8 +72,7 @@ export const previewTts = async (data: {
|
||||
}
|
||||
|
||||
export const getLipsyncJob = async (id: string): Promise<LipsyncJob> => {
|
||||
// MuseTalk 渲染 8s 视频约 54s + 排队时间,给足 5 分钟超时避免单次轮询 AxiosError 中断
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 300_000 })
|
||||
const response = await apiClient.get<LipsyncJob>(`/lipsync/jobs/${id}`, { timeout: 60000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
@@ -92,10 +91,7 @@ export const submitRender = async (data: {
|
||||
}
|
||||
|
||||
export const getRenderJob = async (jobId: string): Promise<RenderJob> => {
|
||||
// 渲染链路(对口型+B-roll+标题+合成+上传)耗时较长,给足 5 分钟超时
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, {
|
||||
timeout: 300_000,
|
||||
})
|
||||
const response = await apiClient.get<RenderJob>(`/ai-avatar/render/${jobId}`, { timeout: 60000 })
|
||||
return response.data
|
||||
}
|
||||
|
||||
|
||||
@@ -1002,11 +1002,6 @@
|
||||
.vv-form-row {
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
.vv-form-hint {
|
||||
font-size: 12px;
|
||||
color: #9ca3af;
|
||||
line-height: 1.4;
|
||||
}
|
||||
@media (max-width: 500px) {
|
||||
.vv-form-grid {
|
||||
grid-template-columns: 1fr;
|
||||
@@ -1096,7 +1091,7 @@
|
||||
line-height: 1.55;
|
||||
}
|
||||
.vv-sb-h {
|
||||
margin: 6px 0 2px;
|
||||
margin: 8px 0 3px;
|
||||
padding: 0;
|
||||
font-size: 14px;
|
||||
font-weight: 700;
|
||||
@@ -1107,20 +1102,48 @@
|
||||
.vv-sb-h:first-child {
|
||||
margin-top: 0;
|
||||
}
|
||||
|
||||
/* 总览 —— 每行一段 */
|
||||
.vv-sb-inline-row {
|
||||
.vv-sb-kv {
|
||||
display: flex;
|
||||
align-items: baseline;
|
||||
flex-wrap: wrap;
|
||||
align-items: flex-start;
|
||||
gap: 4px;
|
||||
font-size: 13px;
|
||||
line-height: 1.7;
|
||||
line-height: 1.55;
|
||||
}
|
||||
.vv-sb-k {
|
||||
flex-shrink: 0;
|
||||
color: #6b7280;
|
||||
font-weight: 500;
|
||||
min-width: 120px;
|
||||
}
|
||||
.vv-sb-kv-ref {
|
||||
align-items: center;
|
||||
}
|
||||
.vv-sb-inline {
|
||||
flex: 1;
|
||||
border: none;
|
||||
background: transparent;
|
||||
color: #1f2937;
|
||||
margin: 2px 0;
|
||||
font-size: 13px;
|
||||
font-family: inherit;
|
||||
padding: 1px 4px;
|
||||
outline: none;
|
||||
border-radius: 4px;
|
||||
}
|
||||
.vv-sb-inline-input {
|
||||
border-bottom: 1px dashed transparent;
|
||||
transition: border-color 0.15s;
|
||||
}
|
||||
.vv-sb-inline-input:hover,
|
||||
.vv-sb-inline-input:focus {
|
||||
border-bottom-color: #7c3aed;
|
||||
background: #f5f0ff;
|
||||
}
|
||||
.vv-sb-inline:disabled {
|
||||
color: #9ca3af;
|
||||
cursor: default;
|
||||
}
|
||||
|
||||
.vv-sb-inline-select {
|
||||
min-width: 80px;
|
||||
min-width: 120px;
|
||||
}
|
||||
.vv-sb-inline-select .ant-select-selector {
|
||||
background: transparent !important;
|
||||
@@ -1135,15 +1158,17 @@
|
||||
line-height: 24px !important;
|
||||
padding-left: 0 !important;
|
||||
}
|
||||
|
||||
/* 段落式 textarea 基础样式(仅编辑态使用) */
|
||||
.vv-sb-doc-ta {
|
||||
flex: 1;
|
||||
min-height: 32px;
|
||||
background: transparent !important;
|
||||
border: 1px dashed transparent !important;
|
||||
padding: 2px 4px !important;
|
||||
font-size: 13px !important;
|
||||
line-height: 1.55 !important;
|
||||
color: #1f2937 !important;
|
||||
border-radius: 6px;
|
||||
border-radius: 4px;
|
||||
resize: vertical;
|
||||
outline: none;
|
||||
transition:
|
||||
border-color 0.15s,
|
||||
background 0.15s;
|
||||
@@ -1153,39 +1178,9 @@
|
||||
border-color: #7c3aed !important;
|
||||
background: #f5f0ff !important;
|
||||
}
|
||||
|
||||
/* 场景与光线 —— 段落样式 */
|
||||
.vv-sb-para {
|
||||
margin: 4px 0;
|
||||
.vv-sb-doc-ta-sm {
|
||||
min-height: 24px;
|
||||
}
|
||||
|
||||
/* 内联编辑 textarea(点击后弹出) */
|
||||
.vv-sb-inline-edit-ta {
|
||||
display: block;
|
||||
width: 100%;
|
||||
margin-top: 4px;
|
||||
min-height: 28px;
|
||||
background: #fafafe !important;
|
||||
border: 1px solid #d8cafc !important;
|
||||
border-radius: 6px !important;
|
||||
padding: 6px 8px !important;
|
||||
font-size: 13px !important;
|
||||
line-height: 1.6 !important;
|
||||
font-family: inherit;
|
||||
color: #1f2937 !important;
|
||||
outline: none;
|
||||
resize: vertical;
|
||||
}
|
||||
.vv-sb-inline-edit-ta-sm {
|
||||
max-width: 120px;
|
||||
}
|
||||
.vv-sb-time-ta {
|
||||
max-width: 160px;
|
||||
font-weight: 600;
|
||||
color: #7c3aed !important;
|
||||
}
|
||||
|
||||
/* 逐镜头 */
|
||||
.vv-sb-doc-shots {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
@@ -1193,60 +1188,25 @@
|
||||
margin-top: 2px;
|
||||
}
|
||||
.vv-sb-doc-shot {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 2px;
|
||||
margin-bottom: 6px;
|
||||
padding-left: 8px;
|
||||
border-left: 2px solid rgba(124, 58, 237, 0.3);
|
||||
}
|
||||
.vv-sb-doc-shot-head {
|
||||
margin-bottom: 1px;
|
||||
}
|
||||
.vv-sb-time-doc {
|
||||
display: block;
|
||||
font-size: 14px;
|
||||
font-weight: 600;
|
||||
color: #7c3aed;
|
||||
margin: 6px 0 2px;
|
||||
cursor: text;
|
||||
}
|
||||
.vv-sb-time-doc:hover {
|
||||
background: rgba(124, 58, 237, 0.06);
|
||||
border-radius: 3px;
|
||||
}
|
||||
|
||||
/* 字段段落 */
|
||||
.vv-sb-field {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
align-items: baseline;
|
||||
margin: 2px 0;
|
||||
font-size: 13px;
|
||||
line-height: 1.6;
|
||||
font-weight: 700;
|
||||
color: #7c3aed;
|
||||
background: transparent;
|
||||
border: none;
|
||||
outline: none;
|
||||
padding: 0 4px 1px;
|
||||
font-family: inherit;
|
||||
border-radius: 4px;
|
||||
}
|
||||
.vv-sb-field-k {
|
||||
color: #1f2937;
|
||||
font-weight: 600;
|
||||
margin-right: 0;
|
||||
white-space: nowrap;
|
||||
}
|
||||
.vv-sb-field-val {
|
||||
color: #374151;
|
||||
cursor: text;
|
||||
border-radius: 3px;
|
||||
padding: 0 2px;
|
||||
transition: background 0.15s;
|
||||
word-break: break-word;
|
||||
flex: 1;
|
||||
min-width: 0;
|
||||
}
|
||||
.vv-sb-field:hover .vv-sb-field-val {
|
||||
background: rgba(124, 58, 237, 0.06);
|
||||
}
|
||||
|
||||
/* 参考图片行 */
|
||||
.vv-sb-ref-row {
|
||||
margin-top: 4px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
flex-wrap: wrap;
|
||||
gap: 4px;
|
||||
.vv-sb-time-doc:focus {
|
||||
background: #f5f0ff;
|
||||
}
|
||||
|
||||
.vv-sb-ref-badge {
|
||||
@@ -1314,10 +1274,6 @@
|
||||
font-size: 12px;
|
||||
}
|
||||
|
||||
/* 标签组段落间距 */
|
||||
.vv-sb-tags-block {
|
||||
margin-top: 4px;
|
||||
}
|
||||
.vv-sb-taglist {
|
||||
display: flex;
|
||||
flex-wrap: wrap;
|
||||
@@ -1380,7 +1336,38 @@
|
||||
background: #f5f0ff;
|
||||
}
|
||||
|
||||
/* 口播稿 —— 复用 vv-sb-field 样式,无额外需求 */
|
||||
.vv-sb-doc-collapse {
|
||||
background: transparent;
|
||||
border: none;
|
||||
color: inherit;
|
||||
font: inherit;
|
||||
cursor: pointer;
|
||||
padding: 0;
|
||||
display: inline-flex;
|
||||
align-items: center;
|
||||
gap: 6px;
|
||||
}
|
||||
.vv-sb-caret {
|
||||
display: inline-block;
|
||||
transition: transform 0.2s;
|
||||
font-size: 10px;
|
||||
color: #9ca3af;
|
||||
}
|
||||
.vv-sb-caret.open {
|
||||
transform: rotate(180deg);
|
||||
}
|
||||
.vv-sb-vo {
|
||||
margin-top: 4px !important;
|
||||
background: #fafafe !important;
|
||||
border-left: 3px solid #7c3aed !important;
|
||||
border-radius: 0 6px 6px 0 !important;
|
||||
padding: 6px 10px !important;
|
||||
min-height: 44px !important;
|
||||
font-size: 13px !important;
|
||||
line-height: 1.7 !important;
|
||||
color: #1f2937 !important;
|
||||
border: none !important;
|
||||
}
|
||||
|
||||
.vv-sb-actions {
|
||||
display: flex;
|
||||
@@ -1469,7 +1456,7 @@
|
||||
}
|
||||
|
||||
/* ── Asset/voice picker modal styles (in page) ───────────── */
|
||||
.vv-modal-mask {
|
||||
.vv-modal {
|
||||
position: fixed;
|
||||
inset: 0;
|
||||
background: rgba(0, 0, 0, 0.45);
|
||||
@@ -1479,59 +1466,21 @@
|
||||
justify-content: center;
|
||||
padding: 20px;
|
||||
}
|
||||
.vv-modal {
|
||||
.vv-modal-body {
|
||||
background: #fff;
|
||||
border: 1px solid #e5e7eb;
|
||||
border-radius: 12px;
|
||||
max-width: 720px;
|
||||
width: 100%;
|
||||
max-height: 80vh;
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
overflow: hidden;
|
||||
position: relative;
|
||||
box-shadow: 0 12px 40px rgba(15, 23, 42, 0.18);
|
||||
}
|
||||
.vv-modal.vv-modal-lg {
|
||||
max-width: 860px;
|
||||
}
|
||||
.vv-modal-head {
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: space-between;
|
||||
padding: 14px 20px;
|
||||
border-bottom: 1px solid #e5e7eb;
|
||||
}
|
||||
.vv-modal-title {
|
||||
font-size: 15px;
|
||||
font-weight: 600;
|
||||
color: #1f2937;
|
||||
}
|
||||
.vv-modal-body {
|
||||
flex: 1 1 auto;
|
||||
overflow-y: auto;
|
||||
padding: 16px 20px;
|
||||
min-height: 0;
|
||||
}
|
||||
.vv-modal-foot {
|
||||
flex-shrink: 0;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: flex-end;
|
||||
gap: 10px;
|
||||
padding: 12px 20px;
|
||||
border-top: 1px solid #e5e7eb;
|
||||
background: #fff;
|
||||
}
|
||||
.vv-modal-foot .vv-btn-primary {
|
||||
width: auto;
|
||||
padding: 8px 18px;
|
||||
}
|
||||
.vv-modal-foot .vv-btn-ghost {
|
||||
padding: 8px 18px;
|
||||
padding: 20px;
|
||||
position: relative;
|
||||
}
|
||||
.vv-modal-close {
|
||||
position: absolute;
|
||||
top: 14px;
|
||||
right: 14px;
|
||||
background: transparent;
|
||||
border: none;
|
||||
color: #6b7280;
|
||||
@@ -1540,11 +1489,6 @@
|
||||
width: 28px;
|
||||
height: 28px;
|
||||
border-radius: 6px;
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
padding: 0;
|
||||
line-height: 1;
|
||||
}
|
||||
.vv-modal-close:hover {
|
||||
color: #ef4444;
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
import React, { useCallback, useEffect, useRef, useState } from "react"
|
||||
import axios from "axios"
|
||||
import {
|
||||
PlusOutlined,
|
||||
CloseOutlined,
|
||||
@@ -40,7 +39,6 @@ import {
|
||||
type ImageAnalysisResult,
|
||||
type CopyResult,
|
||||
type ShotScript,
|
||||
type ViralVideoModel,
|
||||
} from "@/api/viral-video/types"
|
||||
import {
|
||||
generateViralVideo,
|
||||
@@ -48,8 +46,6 @@ import {
|
||||
analyzeViralImages,
|
||||
generateViralCopy,
|
||||
confirmViralCopy,
|
||||
estimateViralVideoCredits,
|
||||
getViralVideoModels,
|
||||
} from "@/api/viral-video"
|
||||
import { useViralVideoPolling } from "./hooks/useViralVideoPolling"
|
||||
import CloneModal from "@/components/voice/CloneModal"
|
||||
@@ -127,9 +123,6 @@ type TabTask = {
|
||||
videoModel: string
|
||||
quality: string
|
||||
upscale: string
|
||||
// 积分预估
|
||||
estimatedCredits: number | null // null=加载中/未发起;浮点数,展示时 toFixed(2)
|
||||
estimatedCreditsError: boolean // true=接口失败显示 --
|
||||
// UI 状态
|
||||
uiStep: UIStep
|
||||
imageAnalysis: ImageAnalysisResult | null
|
||||
@@ -219,46 +212,10 @@ const RATIOS = [
|
||||
{ v: "16:9", label: "16:9 横屏(B站/YouTube)" },
|
||||
{ v: "1:1", label: "1:1 方形(小红书)" },
|
||||
]
|
||||
/** 兜底模型列表(接口未返回时使用,字段与 ViralVideoModel 对齐;后端返回后自动覆盖) */
|
||||
const FALLBACK_VIDEO_MODELS: ViralVideoModel[] = [
|
||||
{
|
||||
key: "seedance-2.5",
|
||||
display_name: "Seedance 2.5 — 最新最强",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p", "1080p"],
|
||||
max_duration: 30,
|
||||
is_default: true,
|
||||
},
|
||||
{
|
||||
key: "seedance-2.0",
|
||||
display_name: "Seedance 2.0 — 正式首选",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p", "1080p"],
|
||||
max_duration: 15,
|
||||
},
|
||||
{
|
||||
key: "seedance-2.0-fast",
|
||||
display_name: "Seedance 2.0 Fast — 快速测试",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p"],
|
||||
max_duration: 15,
|
||||
},
|
||||
{
|
||||
key: "seedance-2.0-mini",
|
||||
display_name: "Seedance 2.0 Mini — 低成本",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p"],
|
||||
max_duration: 15,
|
||||
},
|
||||
{
|
||||
key: "wan-3.0",
|
||||
display_name: "Wan 3.0 — 通义万相全能",
|
||||
supports_audio: true,
|
||||
supported_resolutions: ["480p", "720p", "1080p"],
|
||||
max_duration: 15,
|
||||
},
|
||||
const MODELS = [
|
||||
{ v: "seedance-2.5", label: "Seedance 2.5(推荐)" },
|
||||
{ v: "seedance-2.0", label: "Seedance 2.0" },
|
||||
]
|
||||
const RESOLUTION_ORDER = ["480p", "720p", "1080p", "4k"]
|
||||
const QUALITY_OPTIONS = [
|
||||
{ v: "480p", label: "480p(快速)" },
|
||||
{ v: "720p", label: "720p(清晰)" },
|
||||
@@ -438,8 +395,6 @@ const emptyTask = (id: string, title: string): TabTask => ({
|
||||
videoModel: "seedance-2.5",
|
||||
quality: "480p",
|
||||
upscale: "default",
|
||||
estimatedCredits: null,
|
||||
estimatedCreditsError: false,
|
||||
uiStep: "step1_upload",
|
||||
imageAnalysis: null,
|
||||
storyboard: null,
|
||||
@@ -481,12 +436,12 @@ const ViralVideoPage: React.FC = () => {
|
||||
const [tasks, setTasks] = useState<TabTask[]>(() => [emptyTask("t1", "生成 1")])
|
||||
const [activeId, setActiveId] = useState<string>("t1")
|
||||
const [cloneModalOpen, setCloneModalOpen] = useState(false)
|
||||
const [voOpen, setVoOpen] = useState(false) // 口播稿折叠
|
||||
const [voiceAssets, setVoiceAssets] = useState<
|
||||
{ id: string; name: string; url: string; desc?: string }[]
|
||||
>([])
|
||||
const [presetVoices, setPresetVoices] = useState<PresetVoice[]>([])
|
||||
const [voicePickerOpen, setVoicePickerOpen] = useState(false)
|
||||
const [videoModels, setVideoModels] = useState<ViralVideoModel[]>(FALLBACK_VIDEO_MODELS)
|
||||
const [assetPicker, setAssetPicker] = useState<{
|
||||
open: boolean
|
||||
kind: "image" | "video" | "voice"
|
||||
@@ -546,43 +501,6 @@ const ViralVideoPage: React.FC = () => {
|
||||
})
|
||||
}, [])
|
||||
|
||||
/* ── 加载视频模型列表 ── */
|
||||
useEffect(() => {
|
||||
getViralVideoModels()
|
||||
.then((list) => {
|
||||
// 防御:确保 list 是数组
|
||||
const safeList = Array.isArray(list) ? list : []
|
||||
if (safeList.length === 0) return
|
||||
setVideoModels(safeList)
|
||||
// 如果当前选中的模型不在返回列表里,切换到默认模型并联动修正参数
|
||||
setTask((t) => {
|
||||
const inList = safeList.some((m) => m && typeof m === "object" && m.key === t.videoModel)
|
||||
if (inList) return t
|
||||
const def =
|
||||
safeList.find((m) => m && typeof m === "object" && m.is_default) || safeList[0]
|
||||
if (!def) return t
|
||||
const defRes = Array.isArray(def.supported_resolutions)
|
||||
? def.supported_resolutions
|
||||
: ["480p", "720p", "1080p"]
|
||||
const sortedRes = defRes.slice().sort((a, b) => {
|
||||
const ra = RESOLUTION_ORDER.indexOf(a)
|
||||
const rb = RESOLUTION_ORDER.indexOf(b)
|
||||
return (rb === -1 ? -1 : rb) - (ra === -1 ? -1 : ra)
|
||||
})
|
||||
const newRes = sortedRes[0] || "480p"
|
||||
const newDur = Math.min(
|
||||
t.duration,
|
||||
typeof def.max_duration === "number" ? def.max_duration : 15,
|
||||
)
|
||||
return { ...t, videoModel: def.key, quality: newRes, duration: newDur }
|
||||
})
|
||||
})
|
||||
.catch(() => {
|
||||
// 接口 404/500 时使用兜底列表,不提示用户
|
||||
})
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
|
||||
/* ── 轮询(job 从 STEP1 开始即存在,贯穿三步) ── */
|
||||
const onPollUpdate = useCallback(
|
||||
(job: ViralVideoJob) => {
|
||||
@@ -647,50 +565,6 @@ const ViralVideoPage: React.FC = () => {
|
||||
)
|
||||
useViralVideoPolling(task.jobId, onPollUpdate)
|
||||
|
||||
/* ── 积分动态预估:STEP2文案生成完成后首次调用;STEP3参数变化时防抖300ms刷新 ── */
|
||||
const creditsTimerRef = useRef<ReturnType<typeof setTimeout> | null>(null)
|
||||
const creditsAbortRef = useRef<number>(0)
|
||||
useEffect(() => {
|
||||
const creditsReady =
|
||||
task.uiStep === "step2_copy_ready" ||
|
||||
task.uiStep === "step3_ready" ||
|
||||
task.uiStep === "step3_generating" ||
|
||||
task.uiStep === "step3_done" ||
|
||||
task.uiStep === "failed"
|
||||
if (!creditsReady) {
|
||||
// 未到可预估阶段:清空状态
|
||||
if (task.estimatedCredits !== null || task.estimatedCreditsError) {
|
||||
setTask({ estimatedCredits: null, estimatedCreditsError: false })
|
||||
}
|
||||
return
|
||||
}
|
||||
// 防抖 300ms
|
||||
if (creditsTimerRef.current) clearTimeout(creditsTimerRef.current)
|
||||
const seq = ++creditsAbortRef.current
|
||||
creditsTimerRef.current = setTimeout(async () => {
|
||||
try {
|
||||
const res = await estimateViralVideoCredits({
|
||||
video_model: task.videoModel,
|
||||
resolution: task.quality,
|
||||
video_ratio: task.videoRatio,
|
||||
duration: task.duration,
|
||||
})
|
||||
if (seq !== creditsAbortRef.current) return
|
||||
setTask({
|
||||
estimatedCredits: res.estimated_credits,
|
||||
estimatedCreditsError: false,
|
||||
})
|
||||
} catch {
|
||||
if (seq !== creditsAbortRef.current) return
|
||||
setTask({ estimatedCredits: null, estimatedCreditsError: true })
|
||||
}
|
||||
}, 300)
|
||||
return () => {
|
||||
if (creditsTimerRef.current) clearTimeout(creditsTimerRef.current)
|
||||
}
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [task.videoModel, task.quality, task.videoRatio, task.duration, task.uiStep])
|
||||
|
||||
/* ── OSS 上传 ── */
|
||||
const uploadFile = useCallback(
|
||||
async (file: File, kind: "image" | "video" | "voice", onProgress?: (pct: number) => void) => {
|
||||
@@ -1045,7 +919,6 @@ const ViralVideoPage: React.FC = () => {
|
||||
""
|
||||
job = await confirmViralCopy(task.jobId, {
|
||||
edited_copy: edited && edited !== originalCopy.trim() ? edited : undefined,
|
||||
video_model: task.videoModel,
|
||||
})
|
||||
} else {
|
||||
// 兜底:走旧 /generate 接口(一次性跑完)
|
||||
@@ -1074,21 +947,8 @@ const ViralVideoPage: React.FC = () => {
|
||||
setTask({ job, jobId: job.id })
|
||||
message.success("视频已提交生成,预计 1-3 分钟…")
|
||||
} catch (err: unknown) {
|
||||
// 402 Payment Required → 积分不足
|
||||
if (axios.isAxiosError(err) && err.response?.status === 402) {
|
||||
const detail =
|
||||
(err.response.data as { detail?: string; message?: string; msg?: string } | undefined)
|
||||
?.detail ||
|
||||
(err.response.data as { detail?: string; message?: string; msg?: string } | undefined)
|
||||
?.message ||
|
||||
"积分不足,请充值"
|
||||
message.error(detail)
|
||||
setTask({ uiStep: "step3_ready", videoError: detail })
|
||||
return
|
||||
}
|
||||
const msg = err instanceof Error ? err.message : "提交失败"
|
||||
message.error(msg)
|
||||
setTask({ uiStep: "failed", videoError: msg })
|
||||
message.error(err instanceof Error ? err.message : "提交失败")
|
||||
setTask({ uiStep: "failed", videoError: err instanceof Error ? err.message : "提交失败" })
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1240,85 +1100,40 @@ const ViralVideoPage: React.FC = () => {
|
||||
}
|
||||
|
||||
/* ── 分镜脚本结果区(编导分镜卡片 UI) ── */
|
||||
const [editingField, setEditingField] = useState<string | null>(null)
|
||||
const editTaRef = useRef<HTMLTextAreaElement | null>(null)
|
||||
|
||||
const updateShot = useCallback((idx: number, patch: Partial<StoryboardShot>) => {
|
||||
setTask((t) => {
|
||||
if (!t.storyboard) return t
|
||||
const shots = t.storyboard.shots.map((s, i) => (i === idx ? { ...s, ...patch } : s))
|
||||
return { ...t, storyboard: { ...t.storyboard, shots } }
|
||||
const updateShot = (idx: number, patch: Partial<StoryboardShot>) => {
|
||||
if (!task.storyboard) return
|
||||
const shots = task.storyboard.shots.map((s, i) => (i === idx ? { ...s, ...patch } : s))
|
||||
setTask({ storyboard: { ...task.storyboard, shots } })
|
||||
}
|
||||
const updateStoryboard = (patch: Partial<Storyboard>) => {
|
||||
if (!task.storyboard) return
|
||||
setTask({ storyboard: { ...task.storyboard, ...patch } })
|
||||
}
|
||||
const updateOverview = (patch: Partial<Storyboard["overview"]>) => {
|
||||
if (!task.storyboard) return
|
||||
setTask({
|
||||
storyboard: { ...task.storyboard, overview: { ...task.storyboard.overview, ...patch } },
|
||||
})
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
const updateStoryboard = useCallback((patch: Partial<Storyboard>) => {
|
||||
setTask((t) => (t.storyboard ? { ...t, storyboard: { ...t.storyboard, ...patch } } : t))
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
const updateOverview = useCallback((patch: Partial<Storyboard["overview"]>) => {
|
||||
setTask((t) =>
|
||||
t.storyboard
|
||||
? {
|
||||
...t,
|
||||
storyboard: { ...t.storyboard, overview: { ...t.storyboard.overview, ...patch } },
|
||||
}
|
||||
: t,
|
||||
)
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
const updateConstraints = useCallback(
|
||||
(list: "hard_constraints" | "negative_prompts", idx: number, v: string) => {
|
||||
setTask((t) => {
|
||||
if (!t.storyboard) return t
|
||||
const arr = [...t.storyboard[list]]
|
||||
arr[idx] = v
|
||||
return { ...t, storyboard: { ...t.storyboard, [list]: arr } }
|
||||
})
|
||||
},
|
||||
[], // eslint-disable-line react-hooks/exhaustive-deps
|
||||
)
|
||||
const addConstraint = useCallback((list: "hard_constraints" | "negative_prompts") => {
|
||||
setTask((t) =>
|
||||
t.storyboard
|
||||
? { ...t, storyboard: { ...t.storyboard, [list]: [...t.storyboard[list], ""] } }
|
||||
: t,
|
||||
)
|
||||
// eslint-disable-next-line react-hooks/exhaustive-deps
|
||||
}, [])
|
||||
const removeConstraint = useCallback(
|
||||
(list: "hard_constraints" | "negative_prompts", idx: number) => {
|
||||
setTask((t) => {
|
||||
if (!t.storyboard) return t
|
||||
const arr = t.storyboard[list].filter((_, i) => i !== idx)
|
||||
return { ...t, storyboard: { ...t.storyboard, [list]: arr } }
|
||||
})
|
||||
},
|
||||
[], // eslint-disable-line react-hooks/exhaustive-deps
|
||||
)
|
||||
|
||||
/* 内联编辑 helper —— 点击文本 → textarea,blur/Enter 保存 */
|
||||
const handleInlineCommit = useCallback(() => {
|
||||
const ta = editTaRef.current
|
||||
if (!ta) return
|
||||
const key = ta.dataset.field
|
||||
if (key) {
|
||||
const val = ta.value
|
||||
if (key.startsWith("shot.")) {
|
||||
const parts = key.split(".")
|
||||
const idx = Number(parts[1])
|
||||
const field = parts.slice(2).join(".")
|
||||
updateShot(idx, { [field]: val } as Partial<StoryboardShot>)
|
||||
} else if (key.startsWith("overview.")) {
|
||||
const field = key.split(".")[1]
|
||||
updateOverview({ [field]: val } as Partial<Storyboard["overview"]>)
|
||||
} else if (key === "scene_and_lighting") {
|
||||
updateStoryboard({ scene_and_lighting: val })
|
||||
} else if (key === "voiceover_script") {
|
||||
updateStoryboard({ voiceover_script: val })
|
||||
}
|
||||
}
|
||||
setEditingField(null)
|
||||
}, [updateShot, updateOverview, updateStoryboard])
|
||||
}
|
||||
const updateConstraints = (
|
||||
list: "hard_constraints" | "negative_prompts",
|
||||
idx: number,
|
||||
v: string,
|
||||
) => {
|
||||
if (!task.storyboard) return
|
||||
const arr = [...task.storyboard[list]]
|
||||
arr[idx] = v
|
||||
setTask({ storyboard: { ...task.storyboard, [list]: arr } })
|
||||
}
|
||||
const addConstraint = (list: "hard_constraints" | "negative_prompts") => {
|
||||
if (!task.storyboard) return
|
||||
setTask({ storyboard: { ...task.storyboard, [list]: [...task.storyboard[list], ""] } })
|
||||
}
|
||||
const removeConstraint = (list: "hard_constraints" | "negative_prompts", idx: number) => {
|
||||
if (!task.storyboard) return
|
||||
const arr = task.storyboard[list].filter((_, i) => i !== idx)
|
||||
setTask({ storyboard: { ...task.storyboard, [list]: arr } })
|
||||
}
|
||||
|
||||
const renderCopyResult = () => {
|
||||
if (task.uiStep === "step2_generating") {
|
||||
@@ -1375,108 +1190,46 @@ const ViralVideoPage: React.FC = () => {
|
||||
<div className="vv-sb-doc">
|
||||
{/* 视频总览 */}
|
||||
<h4 className="vv-sb-h">视频总览</h4>
|
||||
<p className="vv-sb-inline-row">
|
||||
<strong className="vv-sb-field-k">整体主题:</strong>
|
||||
{editingField === "overview.theme" ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sb.overview.theme}
|
||||
data-field="overview.theme"
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField("overview.theme")
|
||||
}}
|
||||
>
|
||||
{sb.overview.theme}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-inline-row">
|
||||
<strong className="vv-sb-field-k">总时长:</strong>
|
||||
{editingField === "overview.total_duration" ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta vv-sb-inline-edit-ta-sm"
|
||||
autoFocus
|
||||
defaultValue={sb.overview.total_duration}
|
||||
data-field="overview.total_duration"
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField("overview.total_duration")
|
||||
}}
|
||||
>
|
||||
{sb.overview.total_duration}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-inline-row">
|
||||
<strong className="vv-sb-field-k">画幅:</strong>
|
||||
<div className="vv-sb-kv">
|
||||
<span className="vv-sb-k">整体主题:</span>
|
||||
<input
|
||||
className="vv-sb-inline vv-sb-inline-input"
|
||||
value={sb.overview.theme}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateOverview({ theme: e.target.value })}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-sb-kv">
|
||||
<span className="vv-sb-k">总时长:</span>
|
||||
<input
|
||||
className="vv-sb-inline vv-sb-inline-input"
|
||||
value={sb.overview.total_duration}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateOverview({ total_duration: e.target.value })}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-sb-kv">
|
||||
<span className="vv-sb-k">画幅:</span>
|
||||
<Select
|
||||
className="vv-select vv-sb-inline-select"
|
||||
style={{ width: 90 }}
|
||||
style={{ width: 200 }}
|
||||
value={sb.overview.aspect_ratio}
|
||||
disabled={locked}
|
||||
onChange={(v) => updateOverview({ aspect_ratio: v })}
|
||||
options={RATIOS.map((r) => ({ value: r.v, label: r.v }))}
|
||||
variant="borderless"
|
||||
/>
|
||||
</p>
|
||||
</div>
|
||||
|
||||
{/* 场景与光线 */}
|
||||
<h4 className="vv-sb-h">场景与光线</h4>
|
||||
<p className="vv-sb-para">
|
||||
{editingField === "scene_and_lighting" ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sb.scene_and_lighting}
|
||||
data-field="scene_and_lighting"
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
placeholder="描述整体场景氛围、光线方向与色温…"
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField("scene_and_lighting")
|
||||
}}
|
||||
>
|
||||
{sb.scene_and_lighting}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta"
|
||||
value={sb.scene_and_lighting}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateStoryboard({ scene_and_lighting: e.target.value })}
|
||||
placeholder="描述整体场景氛围、光线方向与色温…"
|
||||
/>
|
||||
|
||||
{/* 逐镜头 */}
|
||||
<h4 className="vv-sb-h">逐秒镜头拆解</h4>
|
||||
@@ -1486,183 +1239,66 @@ const ViralVideoPage: React.FC = () => {
|
||||
typeof sh.reference_image_index === "number"
|
||||
? task.images[sh.reference_image_index]
|
||||
: undefined
|
||||
const shotKey = (field: string) => `shot.${idx}.${field}`
|
||||
return (
|
||||
<div key={idx} className="vv-sb-doc-shot">
|
||||
{editingField === shotKey("time_range") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta vv-sb-time-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.time_range}
|
||||
data-field={shotKey("time_range")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
<div className="vv-sb-doc-shot-head">
|
||||
<input
|
||||
className="vv-sb-time-doc"
|
||||
value={sh.time_range}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { time_range: e.target.value })}
|
||||
placeholder="0-3秒"
|
||||
/>
|
||||
) : (
|
||||
<strong
|
||||
className="vv-sb-time-doc"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("time_range"))
|
||||
}}
|
||||
>
|
||||
{sh.time_range}
|
||||
</strong>
|
||||
)}
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">景别/角度与运镜:</strong>
|
||||
{editingField === shotKey("shot_type_angle_movement") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.shot_type_angle_movement}
|
||||
data-field={shotKey("shot_type_angle_movement")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("shot_type_angle_movement"))
|
||||
}}
|
||||
>
|
||||
{sh.shot_type_angle_movement}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">场景与对白:</strong>
|
||||
{editingField === shotKey("scene_and_dialogue") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.scene_and_dialogue}
|
||||
data-field={shotKey("scene_and_dialogue")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("scene_and_dialogue"))
|
||||
}}
|
||||
>
|
||||
{sh.scene_and_dialogue}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">动作与真人细节:</strong>
|
||||
{editingField === shotKey("action_details") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.action_details}
|
||||
data-field={shotKey("action_details")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("action_details"))
|
||||
}}
|
||||
>
|
||||
{sh.action_details}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">音效/BGM:</strong>
|
||||
{editingField === shotKey("audio_bgm") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.audio_bgm}
|
||||
data-field={shotKey("audio_bgm")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("audio_bgm"))
|
||||
}}
|
||||
>
|
||||
{sh.audio_bgm}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">转场:</strong>
|
||||
{editingField === shotKey("transition") ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sh.transition}
|
||||
data-field={shotKey("transition")}
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField(shotKey("transition"))
|
||||
}}
|
||||
>
|
||||
{sh.transition}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
<p className="vv-sb-field vv-sb-ref-row">
|
||||
<strong className="vv-sb-field-k">参考图片:</strong>
|
||||
</div>
|
||||
<div className="vv-sb-kv">
|
||||
<span className="vv-sb-k">景别/角度与运镜:</span>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-sm"
|
||||
value={sh.shot_type_angle_movement}
|
||||
disabled={locked}
|
||||
onChange={(e) =>
|
||||
updateShot(idx, { shot_type_angle_movement: e.target.value })
|
||||
}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-sb-kv">
|
||||
<span className="vv-sb-k">场景与对白:</span>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta"
|
||||
value={sh.scene_and_dialogue}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { scene_and_dialogue: e.target.value })}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-sb-kv">
|
||||
<span className="vv-sb-k">动作与真人细节:</span>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-sm"
|
||||
value={sh.action_details}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { action_details: e.target.value })}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-sb-kv">
|
||||
<span className="vv-sb-k">音效/BGM:</span>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-sm"
|
||||
value={sh.audio_bgm}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { audio_bgm: e.target.value })}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-sb-kv">
|
||||
<span className="vv-sb-k">转场:</span>
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-doc-ta-sm"
|
||||
value={sh.transition}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateShot(idx, { transition: e.target.value })}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-sb-kv vv-sb-kv-ref">
|
||||
<span className="vv-sb-k">参考图片:</span>
|
||||
{refImg?.url ? (
|
||||
<span className="vv-sb-ref-badge">
|
||||
<img src={refImg.url} alt={`ref-${idx}`} className="vv-sb-ref-thumb" />
|
||||
@@ -1692,7 +1328,7 @@ const ViralVideoPage: React.FC = () => {
|
||||
) : (
|
||||
<span className="vv-sb-muted">未指定</span>
|
||||
)}
|
||||
</p>
|
||||
</div>
|
||||
</div>
|
||||
)
|
||||
})}
|
||||
@@ -1700,100 +1336,79 @@ const ViralVideoPage: React.FC = () => {
|
||||
|
||||
{/* 硬性约束 */}
|
||||
<h4 className="vv-sb-h">硬性约束</h4>
|
||||
<div className="vv-sb-tags-block">
|
||||
<div className="vv-sb-taglist">
|
||||
{sb.hard_constraints.map((c, i) => (
|
||||
<span key={i} className="vv-sb-tag vv-sb-tag-hard">
|
||||
<input
|
||||
className="vv-sb-tag-input"
|
||||
value={c}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateConstraints("hard_constraints", i, e.target.value)}
|
||||
/>
|
||||
{!locked && (
|
||||
<button
|
||||
className="vv-sb-tag-del"
|
||||
onClick={() => removeConstraint("hard_constraints", i)}
|
||||
aria-label="删除"
|
||||
>
|
||||
×
|
||||
</button>
|
||||
)}
|
||||
</span>
|
||||
))}
|
||||
{!locked && (
|
||||
<button className="vv-sb-tag-add" onClick={() => addConstraint("hard_constraints")}>
|
||||
+ 添加
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
<div className="vv-sb-taglist">
|
||||
{sb.hard_constraints.map((c, i) => (
|
||||
<span key={i} className="vv-sb-tag vv-sb-tag-hard">
|
||||
<input
|
||||
className="vv-sb-tag-input"
|
||||
value={c}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateConstraints("hard_constraints", i, e.target.value)}
|
||||
/>
|
||||
{!locked && (
|
||||
<button
|
||||
className="vv-sb-tag-del"
|
||||
onClick={() => removeConstraint("hard_constraints", i)}
|
||||
aria-label="删除"
|
||||
>
|
||||
×
|
||||
</button>
|
||||
)}
|
||||
</span>
|
||||
))}
|
||||
{!locked && (
|
||||
<button className="vv-sb-tag-add" onClick={() => addConstraint("hard_constraints")}>
|
||||
+ 添加
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 负面提示词 */}
|
||||
<h4 className="vv-sb-h">负面提示词</h4>
|
||||
<div className="vv-sb-tags-block">
|
||||
<div className="vv-sb-taglist">
|
||||
{sb.negative_prompts.map((c, i) => (
|
||||
<span key={i} className="vv-sb-tag vv-sb-tag-neg">
|
||||
<input
|
||||
className="vv-sb-tag-input"
|
||||
value={c}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateConstraints("negative_prompts", i, e.target.value)}
|
||||
/>
|
||||
{!locked && (
|
||||
<button
|
||||
className="vv-sb-tag-del"
|
||||
onClick={() => removeConstraint("negative_prompts", i)}
|
||||
aria-label="删除"
|
||||
>
|
||||
×
|
||||
</button>
|
||||
)}
|
||||
</span>
|
||||
))}
|
||||
{!locked && (
|
||||
<button className="vv-sb-tag-add" onClick={() => addConstraint("negative_prompts")}>
|
||||
+ 添加
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
<div className="vv-sb-taglist">
|
||||
{sb.negative_prompts.map((c, i) => (
|
||||
<span key={i} className="vv-sb-tag vv-sb-tag-neg">
|
||||
<input
|
||||
className="vv-sb-tag-input"
|
||||
value={c}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateConstraints("negative_prompts", i, e.target.value)}
|
||||
/>
|
||||
{!locked && (
|
||||
<button
|
||||
className="vv-sb-tag-del"
|
||||
onClick={() => removeConstraint("negative_prompts", i)}
|
||||
aria-label="删除"
|
||||
>
|
||||
×
|
||||
</button>
|
||||
)}
|
||||
</span>
|
||||
))}
|
||||
{!locked && (
|
||||
<button className="vv-sb-tag-add" onClick={() => addConstraint("negative_prompts")}>
|
||||
+ 添加
|
||||
</button>
|
||||
)}
|
||||
</div>
|
||||
|
||||
{/* 完整口播稿 */}
|
||||
<p className="vv-sb-field">
|
||||
<strong className="vv-sb-field-k">
|
||||
<SoundOutlined style={{ color: "#7c3aed", marginRight: 4 }} />
|
||||
口播稿:
|
||||
</strong>
|
||||
{editingField === "voiceover_script" ? (
|
||||
<textarea
|
||||
ref={editTaRef}
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-inline-edit-ta"
|
||||
autoFocus
|
||||
defaultValue={sb.voiceover_script}
|
||||
data-field="voiceover_script"
|
||||
onBlur={handleInlineCommit}
|
||||
onKeyDown={(e) => {
|
||||
if (e.key === "Enter" && !e.shiftKey) {
|
||||
e.preventDefault()
|
||||
handleInlineCommit()
|
||||
}
|
||||
if (e.key === "Escape") setEditingField(null)
|
||||
}}
|
||||
placeholder="AI 合成配音用的完整口播稿…"
|
||||
/>
|
||||
) : (
|
||||
<span
|
||||
className="vv-sb-field-val"
|
||||
onClick={() => {
|
||||
if (!locked) setEditingField("voiceover_script")
|
||||
}}
|
||||
>
|
||||
{sb.voiceover_script}
|
||||
</span>
|
||||
)}
|
||||
</p>
|
||||
{/* 完整口播稿(可折叠) */}
|
||||
<h4 className="vv-sb-h">
|
||||
<button className="vv-sb-doc-collapse" onClick={() => setVoOpen(!voOpen)} type="button">
|
||||
<SoundOutlined style={{ color: "#7c3aed", marginRight: 6 }} />
|
||||
完整口播稿(TTS 合成使用)
|
||||
<span className={`vv-sb-caret ${voOpen ? "open" : ""}`}>▾</span>
|
||||
</button>
|
||||
</h4>
|
||||
{voOpen && (
|
||||
<textarea
|
||||
className="vv-textarea vv-sb-doc-ta vv-sb-vo"
|
||||
value={sb.voiceover_script}
|
||||
disabled={locked}
|
||||
onChange={(e) => updateStoryboard({ voiceover_script: e.target.value })}
|
||||
placeholder="AI 合成配音用的完整口播稿…"
|
||||
/>
|
||||
)}
|
||||
</div>
|
||||
|
||||
<div className="vv-sb-actions">
|
||||
@@ -2388,36 +2003,8 @@ const ViralVideoPage: React.FC = () => {
|
||||
className="vv-select vv-select-step3"
|
||||
style={{ width: "100%" }}
|
||||
value={task.videoModel}
|
||||
onChange={(v) => {
|
||||
const m = videoModels.find((x) => x && x.key === v)
|
||||
if (!m) {
|
||||
setTask({ videoModel: v })
|
||||
return
|
||||
}
|
||||
// 切换模型时自动联动修正分辨率/时长(防御:确保 supported_resolutions 为数组)
|
||||
const mRes = Array.isArray(m.supported_resolutions)
|
||||
? m.supported_resolutions
|
||||
: ["480p", "720p", "1080p"]
|
||||
const sortedRes = mRes.slice().sort((a, b) => {
|
||||
const ra = RESOLUTION_ORDER.indexOf(a)
|
||||
const rb = RESOLUTION_ORDER.indexOf(b)
|
||||
return (rb === -1 ? -1 : rb) - (ra === -1 ? -1 : ra)
|
||||
})
|
||||
const newRes = mRes.includes(task.quality)
|
||||
? task.quality
|
||||
: sortedRes[0] || "480p"
|
||||
const newDur = Math.min(
|
||||
task.duration,
|
||||
typeof m.max_duration === "number" ? m.max_duration : 15,
|
||||
)
|
||||
setTask({ videoModel: v, quality: newRes, duration: newDur })
|
||||
}}
|
||||
options={videoModels
|
||||
.filter((m) => m && typeof m === "object" && m.key)
|
||||
.map((m) => ({
|
||||
value: m.key,
|
||||
label: m.display_name || m.key,
|
||||
}))}
|
||||
onChange={(v) => setTask({ videoModel: v })}
|
||||
options={MODELS.map((m) => ({ value: m.v, label: m.label }))}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-form-row">
|
||||
@@ -2427,20 +2014,7 @@ const ViralVideoPage: React.FC = () => {
|
||||
style={{ width: "100%" }}
|
||||
value={task.quality}
|
||||
onChange={(v) => setTask({ quality: v })}
|
||||
options={(() => {
|
||||
const curM = videoModels.find((m) => m && m.key === task.videoModel)
|
||||
const supported = Array.isArray(curM?.supported_resolutions)
|
||||
? curM!.supported_resolutions
|
||||
: null
|
||||
const pool =
|
||||
supported && supported.length > 0
|
||||
? supported
|
||||
: QUALITY_OPTIONS.map((o) => o.v)
|
||||
return QUALITY_OPTIONS.filter((o) => pool.includes(o.v)).map((o) => ({
|
||||
value: o.v,
|
||||
label: o.label,
|
||||
}))
|
||||
})()}
|
||||
options={QUALITY_OPTIONS.map((m) => ({ value: m.v, label: m.label }))}
|
||||
/>
|
||||
</div>
|
||||
<div className="vv-form-row">
|
||||
@@ -2457,21 +2031,13 @@ const ViralVideoPage: React.FC = () => {
|
||||
</div>
|
||||
</div>
|
||||
|
||||
{step3Enabled && (
|
||||
<div className="vv-credits">
|
||||
<span>
|
||||
<SoundOutlined style={{ marginRight: 4 }} />
|
||||
预计消耗
|
||||
</span>
|
||||
<strong>
|
||||
{task.estimatedCreditsError
|
||||
? "-- 积分"
|
||||
: task.estimatedCredits === null
|
||||
? "… 积分"
|
||||
: `${task.estimatedCredits.toFixed(2)} 积分`}
|
||||
</strong>
|
||||
</div>
|
||||
)}
|
||||
<div className="vv-credits">
|
||||
<span>
|
||||
<SoundOutlined style={{ marginRight: 4 }} />
|
||||
预计消耗
|
||||
</span>
|
||||
<strong>50 积分</strong>
|
||||
</div>
|
||||
|
||||
{/* 视频预览区 */}
|
||||
<div className="vv-section">
|
||||
|
||||
@@ -245,11 +245,12 @@ export default function AssetPickerModal({
|
||||
</div>
|
||||
{multiple && (
|
||||
<div className="vv-modal-foot">
|
||||
<button className="vv-btn vv-btn-ghost" onClick={onClose}>
|
||||
<button className="vv-btn vv-btn-ghost vv-btn-sm" onClick={onClose}>
|
||||
取消
|
||||
</button>
|
||||
<button
|
||||
className="vv-btn vv-btn-primary"
|
||||
style={{ width: "auto", marginTop: 0, padding: "8px 18px" }}
|
||||
onClick={handleConfirm}
|
||||
disabled={picked.size === 0}
|
||||
>
|
||||
|
||||
@@ -27,7 +27,6 @@ import threading
|
||||
import time
|
||||
from concurrent.futures import ThreadPoolExecutor, as_completed
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from celery import Task, shared_task
|
||||
from celery.exceptions import Retry
|
||||
@@ -38,6 +37,7 @@ from packages.adapters.sqlalchemy_impl.viral_video_repository import (
|
||||
SQLAlchemyViralVideoJobRepository,
|
||||
)
|
||||
from packages.domain.viral_video import (
|
||||
CREDITS_VIRAL_VIDEO_COST,
|
||||
STAGE_LABELS,
|
||||
ViralVideoJob,
|
||||
ViralVideoStage,
|
||||
@@ -629,7 +629,7 @@ _SCRIPT_GENERATION_PROMPT = """你是资深短视频导演,为 Seedance 2.5(
|
||||
6. hard_constraints/negative_prompts 保留默认项可追加,不要删减。
|
||||
7. voiceover_script 为纯口播文本(无标记/括号/前缀),{duration}秒约{approx_chars}字。
|
||||
8. 严格按上方「爆款结构」的节奏/段落顺序编排(钩子/痛点/反转/案例/行动号召与结构对齐)。
|
||||
9. 输出前自检:口播对白禁止错别字和语病,**严禁使用"很近",正确用词是"最近"**(指"最近一段时间/最近在用",绝不能写成"很近");其他同音字、形近字错误一律修正。
|
||||
9. 输出前自检:口播对白禁止错别字和语病(特别注意"很/最"等常见误用),同音字错误一律修正。
|
||||
10. 必须使用产品信息中真实的品牌、品名和外观特征,不要编造与产品无关的内容。"""
|
||||
|
||||
|
||||
@@ -647,7 +647,7 @@ def _build_products_summary(image_analysis: dict) -> str:
|
||||
# 优先 VLM 生成的 summary 段(自然语言,给编导模型看效果最好)
|
||||
summary = (p.get("summary") or "").strip()
|
||||
if summary and len(summary) >= 30:
|
||||
lines.append(f"- 图{i + 1} {name}:{summary}")
|
||||
lines.append(f"- 图{i+1} {name}:{summary}")
|
||||
continue
|
||||
# 结构化字段兜底
|
||||
brand = p.get("brand") or ""
|
||||
@@ -669,7 +669,7 @@ def _build_products_summary(image_analysis: dict) -> str:
|
||||
feats = p.get("key_features") or p.get("features") or []
|
||||
sellings = p.get("selling_points") or []
|
||||
scenes = p.get("suitable_scenes") or []
|
||||
parts = [f"图{i + 1} {name}"]
|
||||
parts = [f"图{i+1} {name}"]
|
||||
if brand and brand not in ("未知", "无法判断"):
|
||||
parts.append(f"品牌={brand}")
|
||||
if cat and cat not in ("无法判断", "非产品图"):
|
||||
@@ -729,22 +729,6 @@ def _safe_json_loads(raw: str | dict | list | None):
|
||||
return None
|
||||
|
||||
|
||||
def _replace_henjin_everywhere(obj: Any) -> Any:
|
||||
"""递归遍历 copy_result 里所有字符串值,把'很近'替换成'最近'。
|
||||
覆盖 overview.theme、scene_and_lighting、voiceover_script、
|
||||
shots[].scene_and_dialogue/action_details/audio_bgm 等所有字段。
|
||||
"""
|
||||
if isinstance(obj, str):
|
||||
if "很近" in obj:
|
||||
return obj.replace("很近", "最近")
|
||||
return obj
|
||||
if isinstance(obj, list):
|
||||
return [_replace_henjin_everywhere(x) for x in obj]
|
||||
if isinstance(obj, dict):
|
||||
return {k: _replace_henjin_everywhere(v) for k, v in obj.items()}
|
||||
return obj
|
||||
|
||||
|
||||
def _fallback_script(job: ViralVideoJob) -> dict:
|
||||
"""脚本生成失败时的兜底脚本(极简但可用)。"""
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
@@ -802,7 +786,7 @@ def _validate_and_normalize_script(raw, job: ViralVideoJob) -> dict:
|
||||
continue
|
||||
shots.append(
|
||||
{
|
||||
"time_range": str(s.get("time_range") or f"{i * 3}-{(i + 1) * 3}秒"),
|
||||
"time_range": str(s.get("time_range") or f"{i*3}-{(i+1)*3}秒"),
|
||||
"shot_type_angle_movement": str(s.get("shot_type_angle_movement") or "中景平视,固定镜头"),
|
||||
"scene_and_dialogue": str(s.get("scene_and_dialogue") or ""),
|
||||
"action_details": str(s.get("action_details") or ""),
|
||||
@@ -878,8 +862,8 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
style_hint = "无"
|
||||
if isinstance(job.style_guide, dict):
|
||||
style_hint = (
|
||||
f"节奏{job.style_guide.get('cut_speed', '')}、转场{job.style_guide.get('transition', '')}、"
|
||||
f"色调{job.style_guide.get('color_grade', '')}、能量{job.style_guide.get('energy', '')}"
|
||||
f"节奏{job.style_guide.get('cut_speed','')}、转场{job.style_guide.get('transition','')}、"
|
||||
f"色调{job.style_guide.get('color_grade','')}、能量{job.style_guide.get('energy','')}"
|
||||
)
|
||||
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
@@ -937,15 +921,8 @@ def _step_script_generation(job: ViralVideoJob, intent: dict, image_analysis: di
|
||||
voiceover_len = len(voiceover)
|
||||
shots_cnt = len((normalized or {}).get("shots") or [])
|
||||
# 判定是否"退化到兜底质量":口播过短(<20字)或镜头数<1;正常的短口播(如15s视频~40字)不视为兜底
|
||||
# v1.6.1 P1修复:递归替换 copy_result 里所有字符串字段的"很近"→"最近"(覆盖 overview/scene_and_lighting/voiceover/shots.* 全部字段)
|
||||
_before_dump = json.dumps(normalized, ensure_ascii=False)
|
||||
if "很近" in _before_dump:
|
||||
logger.warning("[爆款视频] 编导脚本含错别字'很近',递归替换为'最近' label=%s", label)
|
||||
normalized = _replace_henjin_everywhere(normalized)
|
||||
voiceover = (normalized or {}).get("voiceover_script") or ""
|
||||
fallback_marker = "我最近在用的好物" in voiceover # _fallback_script 的特征串
|
||||
has_typo_henjin = "很近" in json.dumps(normalized, ensure_ascii=False) # 递归检查仍有"很近"视为不合格
|
||||
is_fallback = fallback_marker or shots_cnt < 1 or voiceover_len < 20 or has_typo_henjin
|
||||
is_fallback = fallback_marker or shots_cnt < 1 or voiceover_len < 20
|
||||
logger.info(
|
||||
"[爆款视频] 编导脚本结果 label=%s voiceover_len=%d shots=%d fallback=%s raw_type=%s",
|
||||
label,
|
||||
@@ -1103,14 +1080,14 @@ def _assemble_seedance_prompt(copy_result: dict, job: ViralVideoJob) -> str:
|
||||
ab = s.get("audio_bgm", "")
|
||||
t = s.get("transition", "")
|
||||
ref = s.get("reference_image_index")
|
||||
lines.append(f"- 镜头{i + 1}({tr}):")
|
||||
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}张产品图")
|
||||
lines.append(f" 参考图片:第{ref+1}张产品图")
|
||||
lines.append("")
|
||||
lines.append("【硬性约束】")
|
||||
for c in hc:
|
||||
@@ -1121,22 +1098,14 @@ def _assemble_seedance_prompt(copy_result: dict, job: ViralVideoJob) -> str:
|
||||
return "\n".join(lines)
|
||||
|
||||
|
||||
def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | None) -> tuple[str, dict | None]:
|
||||
"""步骤 6: v1.6 单次 Seedance 生成(不再分段/拼接)。
|
||||
|
||||
返回 (本地视频路径, usage dict|None)。失败抛异常。
|
||||
"""
|
||||
from packages.domain.points_rules import get_viral_video_model_config
|
||||
def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | None) -> str:
|
||||
"""步骤 6: v1.6 单次 Seedance 生成(不再分段/拼接)。"""
|
||||
from packages.shared.ai_service import call_video_generation
|
||||
|
||||
prompt = _assemble_seedance_prompt(copy_result, job)
|
||||
dur = max(5, min(30, int(getattr(job, "duration", 15) or 15)))
|
||||
ratio = getattr(job, "video_ratio", None) or "9:16"
|
||||
model = getattr(job, "video_model", "") or None
|
||||
resolution = getattr(job, "video_resolution", "720p") or "720p"
|
||||
# 按模型配置决定是否开启音频生成(#2159 多模型支持)
|
||||
_mcfg = get_viral_video_model_config(model)
|
||||
gen_audio = bool(_mcfg.get("supports_audio", True))
|
||||
|
||||
# reference_audios: TTS 音频驱动口型
|
||||
ref_audios = [tts_audio_url] if tts_audio_url else []
|
||||
@@ -1149,12 +1118,10 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
|
||||
tmpdir = Path(tempfile.mkdtemp(prefix=f"viral_{job.id}_"))
|
||||
logger.info(
|
||||
"[爆款视频] 开始单次视频生成 dur=%ds ratio=%s model=%s provider=%s gen_audio=%s ref_imgs=%d ref_audios=%d ref_videos=%d tmpdir=%s",
|
||||
"[爆款视频] 开始单次 Seedance 生成 dur=%ds ratio=%s model=%s ref_imgs=%d ref_audios=%d ref_videos=%d tmpdir=%s",
|
||||
dur,
|
||||
ratio if not first_image else "(follow-image)",
|
||||
model or "default",
|
||||
_mcfg.get("provider", "doubao"),
|
||||
gen_audio,
|
||||
len(rest_images) + (1 if first_image else 0),
|
||||
len(ref_audios),
|
||||
len(ref_videos),
|
||||
@@ -1162,67 +1129,23 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non
|
||||
)
|
||||
logger.info("[爆款视频] Seedance prompt (前300字): %s", prompt[:300])
|
||||
|
||||
# 第一次调用:带参考图/首帧/音频/参考视频
|
||||
result = call_video_generation(
|
||||
video_path = call_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=first_image,
|
||||
duration=dur,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
resolution="720p",
|
||||
output_dir=str(tmpdir),
|
||||
model=model,
|
||||
generate_audio=gen_audio, # 按模型能力:有声模型走原生音画同生;Wan 等需后配 TTS
|
||||
generate_audio=True, # Seedance 原生生成环境音效/BGM;口型由 reference_audios 的 TTS 驱动
|
||||
reference_images=rest_images,
|
||||
reference_audios=ref_audios,
|
||||
reference_videos=ref_videos,
|
||||
)
|
||||
|
||||
# #2169: 真人/肖像拦截已由 ai_client 内部自动切即梦(jimeng-3.0)通道重试——
|
||||
# 保留首帧图、不走"去掉参考图纯 t2v 降级"(用户明确要求按参考照片生成)。
|
||||
# 即梦也失败或非拦截类错误时,直接抛错给上层展示用户友好提示。
|
||||
def _check_and_reraise(result):
|
||||
if result and isinstance(result, dict):
|
||||
return result
|
||||
from packages.shared.ai_service import get_last_video_error
|
||||
|
||||
err = get_last_video_error() or {}
|
||||
user_msg = err.get("user_message") or ""
|
||||
detail = err.get("detail") or ""
|
||||
err_code = err.get("error_code") or "unknown"
|
||||
status_code = err.get("status_code", 0)
|
||||
err_provider = err.get("provider") or _mcfg.get("provider", "doubao")
|
||||
err_msg = user_msg or f"视频生成失败({err_provider} status={status_code} code={err_code})"
|
||||
logger.error(
|
||||
"[爆款视频] 视频生成失败: provider=%s model=%s code=%s status=%s user_msg=%s detail=%s",
|
||||
err_provider,
|
||||
model or "default",
|
||||
err_code,
|
||||
status_code,
|
||||
user_msg,
|
||||
(detail or "")[:500],
|
||||
)
|
||||
raise RuntimeError(err_msg)
|
||||
|
||||
if not result or not isinstance(result, dict):
|
||||
_check_and_reraise(result)
|
||||
video_path = result.get("video_path") or ""
|
||||
usage = result.get("usage")
|
||||
if not video_path or not Path(video_path).exists() or Path(video_path).stat().st_size == 0:
|
||||
raise RuntimeError("视频生成失败:返回空文件或路径不存在")
|
||||
# #2169: 如果实际走了即梦兜底(真人拦截→jimeng),更新 job.video_model 让积分结算用 jimeng-3.0 价格
|
||||
if isinstance(usage, dict):
|
||||
actual_provider = usage.get("provider")
|
||||
actual_model_key = usage.get("model_key")
|
||||
if actual_provider == "jimeng" and actual_model_key:
|
||||
logger.info(
|
||||
"[爆款视频] 实际通过即梦通道生成(原model=%s),更新video_model=%s 用于积分结算",
|
||||
job.video_model,
|
||||
actual_model_key,
|
||||
)
|
||||
job.video_model = actual_model_key
|
||||
size = Path(video_path).stat().st_size
|
||||
logger.info("[爆款视频] 单次生成完成: path=%s size=%d usage=%s", video_path, size, usage)
|
||||
return str(video_path), (usage if isinstance(usage, dict) else None)
|
||||
raise RuntimeError("Seedance 视频生成失败:返回空文件或路径不存在")
|
||||
logger.info("[爆款视频] Seedance 单次生成完成: %s size=%d", video_path, Path(video_path).stat().st_size)
|
||||
return str(video_path)
|
||||
|
||||
|
||||
def _step_upload(job: ViralVideoJob, video_path: str) -> str:
|
||||
@@ -1384,36 +1307,16 @@ def run_video_style_analysis(self: Task, job_id: str) -> dict:
|
||||
|
||||
|
||||
def _mark_failed_and_notify(job_id: str, session, repo, job, err_msg: str, stage: str = "") -> None:
|
||||
"""标记任务失败并通知。若传入的 session 已失效(因前面异常导致 rollback 状态),
|
||||
会自动 fallback 到新建 SessionLocal 重新标记,确保状态一定落库。"""
|
||||
try:
|
||||
# 尝试用传入的 session 标记
|
||||
marked = False
|
||||
if job is not None and not job.is_terminal and session is not None:
|
||||
try:
|
||||
job.mark_failed(err_msg)
|
||||
_save_job(repo, job, session)
|
||||
marked = True
|
||||
except Exception as se:
|
||||
logger.warning("[爆款视频] 用原 session 标记失败失败,fallback 新session: %s", se)
|
||||
try:
|
||||
session.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
if not marked:
|
||||
# fallback:新建独立 session 重新标记(保证状态一定落库)
|
||||
ssn = SessionLocal()
|
||||
try:
|
||||
r = SQLAlchemyViralVideoJobRepository(ssn)
|
||||
j = r.get(job_id)
|
||||
if j is not None and not j.is_terminal:
|
||||
j.mark_failed(err_msg)
|
||||
r.update(j)
|
||||
ssn.commit()
|
||||
finally:
|
||||
ssn.close()
|
||||
if session is None:
|
||||
session = SessionLocal()
|
||||
repo = SQLAlchemyViralVideoJobRepository(session)
|
||||
job = repo.get(job_id)
|
||||
if job is not None and not job.is_terminal:
|
||||
job.mark_failed(err_msg)
|
||||
_save_job(repo, job, session)
|
||||
except Exception as inner:
|
||||
logger.warning("[爆款视频] 标记失败状态时出错(最终fallback也失败): %s", inner, exc_info=True)
|
||||
logger.warning("[爆款视频] 标记失败状态时出错: %s", inner)
|
||||
_emit_progress(
|
||||
job_id,
|
||||
stage,
|
||||
@@ -1640,107 +1543,6 @@ def _quick_compliance_blacklist_check(copy_result: dict) -> None:
|
||||
copy_result[k] = copy_result[k].replace(bk, bv)
|
||||
|
||||
|
||||
def _try_refund_viral_video(job: ViralVideoJob) -> None:
|
||||
"""爆款视频生成失败:若已预扣积分则全额退款。"""
|
||||
try:
|
||||
from packages.shared import get_shared_settings
|
||||
|
||||
_s = get_shared_settings()
|
||||
if not _s.points_enabled:
|
||||
return
|
||||
prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
if prepaid <= 0:
|
||||
return
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
svc = PointsService()
|
||||
# 使用独立 session(避免污染外层事务)
|
||||
ssn = SessionLocal()
|
||||
try:
|
||||
svc.refund_viral_video(
|
||||
job.user_id,
|
||||
prepaid,
|
||||
getattr(job, "credits_transaction_id", "") or "",
|
||||
ssn,
|
||||
)
|
||||
job.credits_prepaid = 0.0
|
||||
finally:
|
||||
ssn.close()
|
||||
except Exception:
|
||||
logger.exception("[爆款视频] 失败退款异常 job_id=%s", job.id)
|
||||
|
||||
|
||||
def _settle_viral_video(job: ViralVideoJob, usage: dict | None) -> None:
|
||||
"""爆款视频生成成功:按实际 usage 结算,多退少补,写 credits_cost。"""
|
||||
try:
|
||||
from packages.shared import get_shared_settings
|
||||
|
||||
_s = get_shared_settings()
|
||||
if not _s.points_enabled:
|
||||
job.credits_cost = 0.0
|
||||
job.credits_prepaid = 0.0
|
||||
return
|
||||
prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
if prepaid <= 0:
|
||||
job.credits_cost = 0.0
|
||||
return
|
||||
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
w, h = resolve_video_dimensions(
|
||||
getattr(job, "video_resolution", "720p") or "720p",
|
||||
getattr(job, "video_ratio", "9:16") or "9:16",
|
||||
)
|
||||
fps = 24
|
||||
duration = int(getattr(job, "duration", 15) or 15)
|
||||
est_tokens = duration * w * h * fps // 1024
|
||||
actual_tokens = None
|
||||
if isinstance(usage, dict):
|
||||
at = usage.get("completion_tokens")
|
||||
if isinstance(at, (int, float)) and at > 0:
|
||||
actual_tokens = int(at)
|
||||
diff_pct = None
|
||||
if actual_tokens and est_tokens > 0:
|
||||
diff_pct = (actual_tokens - est_tokens) * 100.0 / est_tokens
|
||||
logger.info(
|
||||
"[爆款视频] tokens估算vs实际 job_id=%s duration=%s %sx%s model=%s est=%s actual=%s diff=%.1f%%",
|
||||
job.id,
|
||||
duration,
|
||||
w,
|
||||
h,
|
||||
getattr(job, "video_model", "seedance-2.5"),
|
||||
est_tokens,
|
||||
actual_tokens,
|
||||
diff_pct if diff_pct is not None else 0.0,
|
||||
)
|
||||
actual_credits = calculate_viral_video_credits(
|
||||
duration,
|
||||
w,
|
||||
h,
|
||||
getattr(job, "video_model", "") or "seedance-2.5",
|
||||
actual_tokens=actual_tokens,
|
||||
)
|
||||
svc = PointsService()
|
||||
ssn = SessionLocal()
|
||||
try:
|
||||
svc.settle_viral_video(
|
||||
job.user_id,
|
||||
prepaid,
|
||||
actual_credits,
|
||||
getattr(job, "credits_transaction_id", "") or "",
|
||||
ssn,
|
||||
)
|
||||
job.credits_cost = actual_credits
|
||||
job.credits_prepaid = 0.0
|
||||
finally:
|
||||
ssn.close()
|
||||
except Exception:
|
||||
logger.exception("[爆款视频] 积分结算异常 job_id=%s", job.id)
|
||||
# 结算异常不阻塞任务完成:保守按预扣值记 credits_cost
|
||||
job.credits_cost = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
job.credits_prepaid = 0.0
|
||||
|
||||
|
||||
def _run_render_pipeline(job_id: str, session, repo, job) -> dict:
|
||||
"""v1.6.1 阶段3:出片前合规审核(LLM 深度)→ TTS → Seedance → Upload → Completed。
|
||||
|
||||
@@ -1782,17 +1584,9 @@ def _run_render_pipeline(job_id: str, session, repo, job) -> dict:
|
||||
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(失败自动退款)
|
||||
# Step 6: 单次 Seedance
|
||||
_set_stage(job, repo, session, ViralVideoStage.RENDERING, "正在生成视频(约1-3分钟)...")
|
||||
video_path = None
|
||||
usage = None
|
||||
try:
|
||||
video_path, usage = _step_render(job, copy_result, tts_url)
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频][阶段3] Seedance 生成失败,触发退款: %s", e, exc_info=True)
|
||||
# 退款
|
||||
_try_refund_viral_video(job)
|
||||
raise
|
||||
video_path = _step_render(job, copy_result, tts_url)
|
||||
_emit_progress(job_id, ViralVideoStage.RENDERING, 92.0, "视频生成完成")
|
||||
|
||||
# Step 7: Upload
|
||||
@@ -1804,8 +1598,7 @@ def _run_render_pipeline(job_id: str, session, repo, job) -> dict:
|
||||
if video_url:
|
||||
_wait_oss_ready(video_url, timeout_sec=10)
|
||||
|
||||
# 积分结算:按实际 tokens 多退少补
|
||||
_settle_viral_video(job, usage)
|
||||
job.credits_cost = CREDITS_VIRAL_VIDEO_COST
|
||||
job.mark_completed(video_url)
|
||||
job.current_stage = ViralVideoStage.UPLOADING
|
||||
job.phase_message = "视频生成完成"
|
||||
@@ -1843,23 +1636,6 @@ def run_viral_video_render(self: Task, job_id: str) -> dict:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error("[爆款视频][阶段3] 异常: %s", e, exc_info=True)
|
||||
# 兜底:任何阶段3异常都尝试退款(_step_render 内部异常已经退过,但 upload 等后续失败也需退)
|
||||
try:
|
||||
if session is not None:
|
||||
job_safe = None
|
||||
try:
|
||||
repo_safe = SQLAlchemyViralVideoJobRepository(session)
|
||||
job_safe = repo_safe.get(job_id)
|
||||
except Exception:
|
||||
pass
|
||||
if job_safe is not None and float(getattr(job_safe, "credits_prepaid", 0) or 0) > 0:
|
||||
_try_refund_viral_video(job_safe)
|
||||
try:
|
||||
repo_safe.update(job_safe)
|
||||
except Exception:
|
||||
pass
|
||||
except Exception:
|
||||
logger.exception("[爆款视频][阶段3] 兜底退款异常")
|
||||
_mark_failed_and_notify(job_id, session, None, None, str(e), ViralVideoStage.RENDERING)
|
||||
return {"ok": False, "job_id": job_id, "error": str(e)}
|
||||
finally:
|
||||
|
||||
@@ -56,7 +56,7 @@ class UserModel(Base):
|
||||
is_member = Column(Boolean, nullable=False, default=False)
|
||||
member_type = Column(String(20), nullable=True)
|
||||
member_expires_at = Column(DateTime, nullable=True)
|
||||
points_balance = Column(Float, nullable=False, default=0)
|
||||
points_balance = Column(Integer, nullable=False, default=0)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
@@ -775,9 +775,9 @@ class PointsAccountModel(Base):
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
user_id = Column(String(36), nullable=False, unique=True, index=True)
|
||||
balance = Column(Float, nullable=False, default=0)
|
||||
total_earned = Column(Float, nullable=False, default=0)
|
||||
total_spent = Column(Float, nullable=False, default=0)
|
||||
balance = Column(Integer, nullable=False, default=0)
|
||||
total_earned = Column(Integer, nullable=False, default=0)
|
||||
total_spent = Column(Integer, nullable=False, default=0)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
|
||||
@@ -792,8 +792,8 @@ class PointsTransactionModel(Base):
|
||||
account_id = Column(String(36), nullable=False, index=True)
|
||||
type = Column(String(20), nullable=False, index=True) # earn / spend / refund
|
||||
source = Column(String(50), nullable=False, index=True)
|
||||
amount = Column(Float, nullable=False)
|
||||
balance_after = Column(Float, nullable=False)
|
||||
amount = Column(Integer, nullable=False)
|
||||
balance_after = Column(Integer, nullable=False)
|
||||
description = Column(String(255), nullable=False, default="")
|
||||
ref_id = Column(String(100), nullable=False, default="")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC))
|
||||
@@ -961,10 +961,7 @@ class ViralVideoJobModel(Base):
|
||||
JSON, nullable=True
|
||||
) # v1.6: 编导脚本结构{overview,scene_and_lighting,shots,hard_constraints,negative_prompts,voiceover_script}
|
||||
result_video_url = Column(String(1000), nullable=False, default="")
|
||||
credits_cost = Column(Float, nullable=False, default=0)
|
||||
video_resolution = Column(String(20), nullable=False, default="720p")
|
||||
credits_prepaid = Column(Float, nullable=False, default=0.0)
|
||||
credits_transaction_id = Column(String(36), nullable=False, default="")
|
||||
credits_cost = Column(Integer, nullable=False, default=0)
|
||||
error_msg = Column(Text, nullable=False, default="")
|
||||
retry_count = Column(Integer, nullable=False, default=0)
|
||||
started_at = Column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
@@ -39,9 +39,6 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
model.phone_verified = user.phone_verified
|
||||
model.binding_completed_at = user.binding_completed_at
|
||||
model.profile_completed = user.profile_completed
|
||||
model.is_member = user.is_member
|
||||
model.member_type = user.member_type
|
||||
model.member_expires_at = user.member_expires_at
|
||||
model.created_at = user.created_at
|
||||
|
||||
self.session.commit()
|
||||
@@ -118,8 +115,5 @@ class SQLAlchemyUserRepository(UserRepository):
|
||||
phone_verified=model.phone_verified or False,
|
||||
binding_completed_at=model.binding_completed_at,
|
||||
profile_completed=model.profile_completed if model.profile_completed is not None else True,
|
||||
is_member=model.is_member if model.is_member is not None else False,
|
||||
member_type=model.member_type,
|
||||
member_expires_at=model.member_expires_at,
|
||||
created_at=model.created_at,
|
||||
)
|
||||
|
||||
@@ -46,10 +46,7 @@ def _to_domain(model: ViralVideoJobModel) -> ViralVideoJob:
|
||||
generated_copy_text=getattr(model, "generated_copy_text", "") or "",
|
||||
copy_result=dict(model.copy_result) if getattr(model, "copy_result", None) else None,
|
||||
result_video_url=model.result_video_url or "",
|
||||
video_resolution=getattr(model, "video_resolution", "720p") or "720p",
|
||||
credits_prepaid=float(getattr(model, "credits_prepaid", 0) or 0),
|
||||
credits_transaction_id=getattr(model, "credits_transaction_id", "") or "",
|
||||
credits_cost=float(model.credits_cost or 0),
|
||||
credits_cost=model.credits_cost or 0,
|
||||
error_msg=model.error_msg or "",
|
||||
retry_count=model.retry_count or 0,
|
||||
started_at=model.started_at,
|
||||
@@ -98,10 +95,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
generated_copy_text=job.generated_copy_text,
|
||||
copy_result=job.copy_result,
|
||||
result_video_url=job.result_video_url,
|
||||
video_resolution=getattr(job, "video_resolution", "720p") or "720p",
|
||||
credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0),
|
||||
credits_transaction_id=getattr(job, "credits_transaction_id", "") or "",
|
||||
credits_cost=float(getattr(job, "credits_cost", 0) or 0),
|
||||
credits_cost=job.credits_cost,
|
||||
error_msg=job.error_msg,
|
||||
retry_count=job.retry_count,
|
||||
started_at=job.started_at,
|
||||
@@ -127,10 +121,7 @@ class SQLAlchemyViralVideoJobRepository:
|
||||
model.generated_copy_text = job.generated_copy_text or ""
|
||||
model.copy_result = job.copy_result
|
||||
model.result_video_url = job.result_video_url
|
||||
model.video_resolution = getattr(job, "video_resolution", "720p") or "720p"
|
||||
model.credits_prepaid = float(getattr(job, "credits_prepaid", 0) or 0)
|
||||
model.credits_transaction_id = getattr(job, "credits_transaction_id", "") or ""
|
||||
model.credits_cost = float(getattr(job, "credits_cost", 0) or 0)
|
||||
model.credits_cost = job.credits_cost
|
||||
model.error_msg = job.error_msg
|
||||
model.retry_count = job.retry_count
|
||||
model.started_at = job.started_at
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
"""应用层:对外展示目录(套餐/积分包)。"""
|
||||
@@ -1,152 +0,0 @@
|
||||
"""读取管理后台配置的会员套餐 / 积分充值包(共享库真实数据)。
|
||||
|
||||
替代旧的硬编码 MEMBERSHIP_PRICES / POINTS_PACKAGES。
|
||||
短 TTL 缓存(30 秒),后台改价/启停后用户端最多 30 秒可见。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
_CACHE_TTL = 30.0
|
||||
_lock = threading.Lock()
|
||||
_cache: dict[str, tuple[float, Any]] = {}
|
||||
|
||||
_QUOTA_LABELS = {
|
||||
"4k": "4K 超清分辨率",
|
||||
"batch_render": "批量渲染",
|
||||
"priority_queue": "优先处理队列",
|
||||
"ai_matting": "AI 智能抠像",
|
||||
"remove_watermark": "去水印",
|
||||
}
|
||||
|
||||
|
||||
def _cached(key: str, loader):
|
||||
now = time.time()
|
||||
hit = _cache.get(key)
|
||||
if hit and now - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
with _lock:
|
||||
hit = _cache.get(key)
|
||||
if hit and time.time() - hit[0] < _CACHE_TTL:
|
||||
return hit[1]
|
||||
value = loader()
|
||||
_cache[key] = (time.time(), value)
|
||||
return value
|
||||
|
||||
|
||||
def _quota_features(quotas: dict[str, Any] | None) -> dict[str, Any]:
|
||||
quotas = quotas or {}
|
||||
features: dict[str, Any] = {}
|
||||
for k, v in quotas.items():
|
||||
if k == "credits_per_month":
|
||||
features["credits_per_month"] = v
|
||||
elif k in _QUOTA_LABELS:
|
||||
features[_QUOTA_LABELS[k]] = v
|
||||
else:
|
||||
features[k] = v
|
||||
return features
|
||||
|
||||
|
||||
def get_membership_plans() -> list[dict[str, Any]]:
|
||||
"""读取 is_enabled=true 的套餐,按年/月周期展开为用户端档位。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT plan_key, name, description, monthly_price, yearly_price,
|
||||
quotas, display_order
|
||||
FROM plans
|
||||
WHERE is_enabled = TRUE
|
||||
ORDER BY display_order NULLS LAST, created_at
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
plans: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
base_features = _quota_features(r.quotas if isinstance(r.quotas, dict) else None)
|
||||
if r.yearly_price and float(r.yearly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "yearly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.yearly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.yearly_price) * 100 / 12)),
|
||||
"duration_days": 365,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
if r.monthly_price and float(r.monthly_price) > 0:
|
||||
plans.append(
|
||||
{
|
||||
"plan_id": r.plan_key,
|
||||
"billing_cycle": "monthly",
|
||||
"name": r.name,
|
||||
"description": r.description,
|
||||
"price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"monthly_price_cents": int(round(float(r.monthly_price) * 100)),
|
||||
"duration_days": 30,
|
||||
"features": dict(base_features),
|
||||
}
|
||||
)
|
||||
return plans
|
||||
|
||||
return _cached("membership_plans", _load)
|
||||
|
||||
|
||||
def get_points_packages() -> list[dict[str, Any]]:
|
||||
"""读取 is_active=true 的积分充值包。"""
|
||||
|
||||
def _load() -> list[dict[str, Any]]:
|
||||
from sqlalchemy import text
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.session import SessionLocal
|
||||
|
||||
if SessionLocal is None:
|
||||
return []
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
rows = session.execute(text("""
|
||||
SELECT package_key, name, price, credits, bonus_credits,
|
||||
is_recommended, description, sort_order
|
||||
FROM credit_packages
|
||||
WHERE is_active = TRUE
|
||||
ORDER BY sort_order NULLS LAST, price
|
||||
""")).fetchall()
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
packages: list[dict[str, Any]] = []
|
||||
for r in rows:
|
||||
total_points = int(r.credits or 0) + int(r.bonus_credits or 0)
|
||||
price_cents = int(round(float(r.price) * 100))
|
||||
unit = (price_cents / 100 / total_points) if total_points else 0
|
||||
packages.append(
|
||||
{
|
||||
"code": r.package_key,
|
||||
"name": r.name,
|
||||
"points": total_points,
|
||||
"bonus_credits": int(r.bonus_credits or 0),
|
||||
"price_cents": price_cents,
|
||||
"unit_price": f"¥{unit:.3f}/积分",
|
||||
"is_recommended": bool(r.is_recommended),
|
||||
"description": r.description,
|
||||
}
|
||||
)
|
||||
return packages
|
||||
|
||||
return _cached("points_packages", _load)
|
||||
+5
-28
@@ -90,42 +90,19 @@ class SharedSettings(BaseSettings):
|
||||
|
||||
# ── 豆包大模型(火山引擎方舟) ────────────────────────────────────────
|
||||
doubao_api_key: str = ""
|
||||
doubao_model: str = "doubao-seed-2-1-pro-260915" # 推理模型(Seed 2.1 Pro,深度思考+多模态;原 seed-1-6 已下线)
|
||||
doubao_fast_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速模型(Seed 2.1 Lite,高 RPM,编导/审核/VLM lite;原 1-5-pro-32k 已 Retiring)
|
||||
)
|
||||
doubao_model: str = "doubao-seed-1-6-250615" # 推理模型(通用兜底)
|
||||
doubao_fast_model: str = "doubao-1-5-pro-32k-250115" # 快速结构化输出模型(编导脚本/意图解析/审核)
|
||||
doubao_base_url: str = "https://ark.cn-beijing.volces.com/api/v3"
|
||||
doubao_timeout: int = 30
|
||||
doubao_max_retries: int = 2
|
||||
doubao_vision_model: str = (
|
||||
"doubao-seed-2-1-pro-260915" # 高精度视觉(Seed 2.1 Pro 原生多模态;原 vision-pro-250328 已下线)
|
||||
)
|
||||
doubao_vision_lite_model: str = (
|
||||
"doubao-seed-2-1-lite-260915" # 快速视觉(Seed 2.1 Lite 原生多模态;原 vision-lite-250315 不可用)
|
||||
)
|
||||
doubao_vision_model: str = "doubao-1-5-vision-pro-250328" # 高精度视觉(备用)
|
||||
doubao_vision_lite_model: str = "doubao-1-5-vision-lite-250315" # 快速视觉(商品识别默认,速度优先)
|
||||
doubao_vision_use_lite: bool = True # viral-video 图片分析默认用 lite 提速
|
||||
doubao_embedding_model: str = "doubao-embedding-vision-251215" # 多模态向量化(原 large-text-240915 已 Retiring)
|
||||
doubao_embedding_model: str = "doubao-embedding-large-text-240915"
|
||||
doubao_video_model: str = "doubao-seedance-2-5-260628"
|
||||
doubao_video_timeout: int = 600 # 视频生成轮询总超时(秒)
|
||||
doubao_video_poll_interval: int = 10 # 轮询间隔(秒)
|
||||
|
||||
# ── DashScope (阿里云百炼 Wan 3.0 等) ─────────────────────────────────
|
||||
dashscope_api_key: str = ""
|
||||
dashscope_base_url: str = "https://dashscope.aliyuncs.com/api/v1"
|
||||
dashscope_video_timeout: int = 900 # Wan 视频任务轮询总超时(秒)
|
||||
dashscope_video_poll_interval: int = 10
|
||||
|
||||
# ── 即梦(Jimeng)视觉 API —— 火山引擎 cvtob ──────────────────────────
|
||||
# #2169: 真人参考图在方舟 Seedance 走 B 端审核会被 50411 拦截,
|
||||
# 即梦走 C 端审核链路,普通真人照片可过审,作为参考图场景兜底通道。
|
||||
# 鉴权:AK/SK V4 签名(Region=cn-north-1, Service=cv)
|
||||
jimeng_ak: str = ""
|
||||
jimeng_sk: str = ""
|
||||
jimeng_base_url: str = "https://visual.volcengineapi.com"
|
||||
jimeng_req_key: str = "jimeng_i2v_first_v30" # 视频3.0 720P 首帧图生视频(1张图,稳定版;1080P 标注下线中)
|
||||
jimeng_video_timeout: int = 600 # 即梦轮询总超时(秒)
|
||||
jimeng_video_poll_interval: int = 5 # 轮询间隔(秒)
|
||||
|
||||
# ── MediaKit (火山引擎 AI 媒体工具) ──────────────────────────────────
|
||||
mediakit_api_key: str = ""
|
||||
mediakit_base_url: str = "https://mediakit.cn-beijing.volces.com/api/v1"
|
||||
|
||||
@@ -60,11 +60,6 @@ class User:
|
||||
# 资料是否已完善(微信新用户首次设置昵称后置 True;邮箱注册默认 True)
|
||||
profile_completed: bool = True
|
||||
|
||||
# 会员字段 (#1895):与 users 表列对应
|
||||
is_member: bool = False
|
||||
member_type: str | None = None
|
||||
member_expires_at: datetime | None = None
|
||||
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
|
||||
@@ -9,9 +9,9 @@ from uuid import uuid4
|
||||
class PointsAccount:
|
||||
id: str
|
||||
user_id: str
|
||||
balance: float = 0.0
|
||||
total_earned: float = 0.0
|
||||
total_spent: float = 0.0
|
||||
balance: int = 0
|
||||
total_earned: int = 0
|
||||
total_spent: int = 0
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
|
||||
|
||||
+54
-349
@@ -1,340 +1,32 @@
|
||||
"""积分消耗规则配置 (#1895)
|
||||
|
||||
v1.6.1: 按产品决策,智能混剪/AI数字人/AI配音/抖音解析/改写/标题/封面 全部免费,
|
||||
仅保留声音克隆合成(voice_clone_synth)的扣点逻辑;声音克隆训练保持免费。
|
||||
爆款视频(viral_video)走动态定价,见本文件 VIRAL_VIDEO_MODEL_PRICES + calculate_viral_video_credits。
|
||||
"""
|
||||
"""积分消耗规则配置 (#1895)"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
|
||||
# ============ 爆款视频动态定价 (#2151) ============
|
||||
# key = (model_id, resolution, has_video_input),单位:
|
||||
# - billing_mode=token: 元/百万tokens(输出)
|
||||
# - billing_mode=per_second: 元/秒(视频时长)
|
||||
VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = {
|
||||
("seedance-2.5", "480p", False): 70.0,
|
||||
("seedance-2.5", "720p", False): 70.0,
|
||||
("seedance-2.5", "1080p", False): 77.0,
|
||||
("seedance-2.5", "480p", True): 42.0,
|
||||
("seedance-2.5", "720p", True): 42.0,
|
||||
("seedance-2.5", "1080p", True): 46.0,
|
||||
("seedance-2.0", "480p", False): 46.0,
|
||||
("seedance-2.0", "720p", False): 46.0,
|
||||
("seedance-2.0", "1080p", False): 51.0,
|
||||
("seedance-2.0", "4k", False): 80.0,
|
||||
("seedance-2.0-fast", "480p", False): 28.0,
|
||||
("seedance-2.0-fast", "720p", False): 28.0,
|
||||
("seedance-2.0-mini", "480p", False): 9.2,
|
||||
("seedance-2.0-mini", "720p", False): 9.2,
|
||||
("wan-3.0", "480p", False): 0.3,
|
||||
("wan-3.0", "720p", False): 0.6,
|
||||
("wan-3.0", "1080p", False): 1.2,
|
||||
# #2169: 即梦(Jimeng)视频3.0 720P 首帧图生视频,0.28 元/秒(C 端审核,真人可过)
|
||||
("jimeng-3.0", "720p", False): 0.28,
|
||||
}
|
||||
|
||||
# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器
|
||||
VIRAL_VIDEO_FIXED_COST = 0.15
|
||||
# 利润系数
|
||||
VIRAL_VIDEO_PROFIT_MULTIPLIER = 1.3
|
||||
# Seedance 输出帧率
|
||||
VIRAL_VIDEO_FPS = 24
|
||||
|
||||
# 分辨率别名映射 -> 标准 key
|
||||
_RESOLUTION_ALIASES: dict[str, str] = {
|
||||
"480p": "480p",
|
||||
"普清": "480p",
|
||||
"default": "480p",
|
||||
"low": "480p",
|
||||
"sd": "480p",
|
||||
"720p": "720p",
|
||||
"高清": "720p",
|
||||
"medium": "720p",
|
||||
"hd": "720p",
|
||||
"1080p": "1080p",
|
||||
"超清": "1080p",
|
||||
"high": "1080p",
|
||||
"ultra": "1080p",
|
||||
"全能": "1080p",
|
||||
"fhd": "1080p",
|
||||
}
|
||||
# 分辨率 -> 短边像素数(p 值代表短边,不是 height)
|
||||
_RESOLUTION_SHORT_SIDE: dict[str, int] = {"480p": 480, "720p": 720, "1080p": 1080, "4k": 2160}
|
||||
_RESOLUTION_ALIASES["4k"] = "4k"
|
||||
_RESOLUTION_ALIASES["2160p"] = "4k"
|
||||
_RESOLUTION_ALIASES["uhd"] = "4k"
|
||||
|
||||
|
||||
def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]:
|
||||
"""把 (resolution, ratio) 解析为 (width, height)。
|
||||
|
||||
resolution 数字代表短边像素数(480p/720p/1080p 等):
|
||||
- 横屏 16:9:短边是 height,width = short * 16/9
|
||||
- 竖屏 9:16:短边是 width,height = short * 16/9
|
||||
- 方屏 1:1:width = height = short
|
||||
"""
|
||||
key = str(resolution or "").strip()
|
||||
key_l = key.lower()
|
||||
res_key = _RESOLUTION_ALIASES.get(key_l) or _RESOLUTION_ALIASES.get(key) or "720p"
|
||||
short = _RESOLUTION_SHORT_SIDE.get(res_key, 720)
|
||||
r = str(ratio or "").strip().lower()
|
||||
if r == "16:9":
|
||||
# 横屏:短边是 height,width 向上取整并对齐偶数
|
||||
w = math.ceil(short * 16 / 9)
|
||||
h = short
|
||||
elif r == "1:1":
|
||||
w, h = short, short
|
||||
else:
|
||||
# 9:16 竖屏(默认):短边是 width,height 向上取整并对齐偶数
|
||||
w = short
|
||||
h = math.ceil(short * 16 / 9)
|
||||
# 对齐到偶数(视频编码要求)
|
||||
w = w + (w % 2)
|
||||
h = h + (h % 2)
|
||||
return int(w), int(h)
|
||||
|
||||
|
||||
# ── 爆款视频多模型元数据 (#2159) ──────────────────────────────────────
|
||||
VIRAL_VIDEO_MODEL_CONFIG: dict[str, dict] = {
|
||||
"seedance-2.5": {
|
||||
"key": "seedance-2.5",
|
||||
"display_name": "Seedance 2.5 — 最新最强",
|
||||
"model_id": "doubao-seedance-2-5-260628",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p", "1080p"],
|
||||
"max_duration": 30,
|
||||
"billing_mode": "token",
|
||||
"is_default": True,
|
||||
},
|
||||
"seedance-2.0": {
|
||||
"key": "seedance-2.0",
|
||||
"display_name": "Seedance 2.0 — 正式首选",
|
||||
"model_id": "doubao-seedance-2-0-260128",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p", "1080p", "4k"],
|
||||
"max_duration": 15,
|
||||
"billing_mode": "token",
|
||||
"is_default": False,
|
||||
},
|
||||
"seedance-2.0-fast": {
|
||||
"key": "seedance-2.0-fast",
|
||||
"display_name": "Seedance 2.0 Fast — 快速低成本",
|
||||
"model_id": "doubao-seedance-2-0-fast-260128",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p"],
|
||||
"max_duration": 15,
|
||||
"billing_mode": "token",
|
||||
"is_default": False,
|
||||
},
|
||||
"seedance-2.0-mini": {
|
||||
"key": "seedance-2.0-mini",
|
||||
"display_name": "Seedance 2.0 Mini — 低成本测试",
|
||||
"model_id": "doubao-seedance-2-0-mini-260615",
|
||||
"provider": "doubao",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p"],
|
||||
"max_duration": 15,
|
||||
"billing_mode": "token",
|
||||
"is_default": False,
|
||||
},
|
||||
"wan-3.0": {
|
||||
"key": "wan-3.0",
|
||||
"display_name": "Wan 3.0 — 通义万相(阿里云)",
|
||||
"model_id": "wan3.0-video",
|
||||
"provider": "dashscope",
|
||||
"supports_audio": True,
|
||||
"supported_resolutions": ["480p", "720p", "1080p"],
|
||||
"max_duration": 30,
|
||||
"billing_mode": "per_second",
|
||||
"is_default": False,
|
||||
},
|
||||
# #2169: 即梦视频3.0(内部兜底通道,方舟 Seedance 返回真人拦截 50411 时自动切到即梦重试,
|
||||
# 不暴露给前端让用户直接选择,但需支持计费结算)
|
||||
"jimeng-3.0": {
|
||||
"key": "jimeng-3.0",
|
||||
"display_name": "即梦3.0 — 真人图生视频(兜底)",
|
||||
"model_id": "jimeng_i2v_first_v30",
|
||||
"provider": "jimeng",
|
||||
"supports_audio": False, # 即梦返回无声视频,音频由后续 ffmpeg 合成 TTS
|
||||
"supported_resolutions": ["720p"],
|
||||
"max_duration": 10, # 即梦 i2v 首帧最长 10s(frames=241)
|
||||
"billing_mode": "per_second",
|
||||
"is_default": False,
|
||||
"_internal_fallback_only": True, # 标记:不对外暴露到模型选择列表
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def get_viral_video_model_config(model_key: str | None) -> dict:
|
||||
"""获取模型配置,未知 key 回落到默认 seedance-2.5。"""
|
||||
key = (model_key or "").strip().lower()
|
||||
if key and key in VIRAL_VIDEO_MODEL_CONFIG:
|
||||
return VIRAL_VIDEO_MODEL_CONFIG[key]
|
||||
return VIRAL_VIDEO_MODEL_CONFIG["seedance-2.5"]
|
||||
|
||||
|
||||
def list_viral_video_models(
|
||||
include_placeholder: bool = False,
|
||||
dashscope_available: bool = False,
|
||||
) -> list[dict]:
|
||||
"""返回前端可用的模型列表(供 GET /api/v1/viral-video/models 端点用)。"""
|
||||
out: list[dict] = []
|
||||
for _k, cfg in VIRAL_VIDEO_MODEL_CONFIG.items():
|
||||
if cfg.get("_placeholder") and not include_placeholder:
|
||||
continue
|
||||
if cfg.get("provider") == "dashscope" and not dashscope_available:
|
||||
continue
|
||||
# #2169: 即梦是内部兜底通道,不在前端模型列表展示
|
||||
if cfg.get("_internal_fallback_only"):
|
||||
continue
|
||||
out.append(
|
||||
{
|
||||
"key": cfg["key"],
|
||||
"display_name": cfg["display_name"],
|
||||
"supports_audio": bool(cfg.get("supports_audio", True)),
|
||||
"supported_resolutions": list(cfg.get("supported_resolutions", ["720p"])),
|
||||
"max_duration": int(cfg.get("max_duration", 15)),
|
||||
"billing_mode": cfg.get("billing_mode", "token"),
|
||||
"is_default": bool(cfg.get("is_default", False)),
|
||||
}
|
||||
)
|
||||
return out
|
||||
|
||||
|
||||
def _match_model_prefix(model: str | None) -> str:
|
||||
"""匹配 model key(支持全部内部别名,未知回落到 seedance-2.5)。
|
||||
|
||||
按 key 长度从长到短匹配,避免 "seedance-2.0-fast" 被 "seedance-2.0" 前缀命中。
|
||||
"""
|
||||
mm = (model or "").strip().lower()
|
||||
for k in sorted(VIRAL_VIDEO_MODEL_CONFIG.keys(), key=len, reverse=True):
|
||||
if mm == k or mm.startswith(k):
|
||||
return k
|
||||
return "seedance-2.5"
|
||||
|
||||
|
||||
def _infer_resolution_key(width: int, height: int) -> str:
|
||||
"""从实际 (width, height) 用短边推断 resolution key。"""
|
||||
short = min(int(width or 720), int(height or 720))
|
||||
if short >= 1900:
|
||||
return "4k"
|
||||
if short >= 1000:
|
||||
return "1080p"
|
||||
if short >= 650:
|
||||
return "720p"
|
||||
return "480p"
|
||||
|
||||
|
||||
def calculate_viral_video_credits_with_breakdown(
|
||||
duration_seconds: int,
|
||||
width: int,
|
||||
height: int,
|
||||
model: str = "seedance-2.5",
|
||||
has_video_input: bool = False,
|
||||
actual_tokens: int | None = None,
|
||||
fps: int = VIRAL_VIDEO_FPS,
|
||||
) -> tuple[float, dict]:
|
||||
"""计算爆款视频所需积分(1 积分 = 1 元),并返回计费公式明细。
|
||||
|
||||
公式:
|
||||
tokens = duration * width * height * fps / 1024
|
||||
video_cost = tokens / 1_000_000 * model_token_price
|
||||
total = round((video_cost + fixed_cost) * profit_multiplier, 2)
|
||||
若传入 actual_tokens 则用它替代计算值。
|
||||
|
||||
Returns:
|
||||
(credits, breakdown) 二元组:
|
||||
- credits: 四舍五入保留两位小数的最终积分
|
||||
- breakdown: dict,包含 tokens / video_cost / fixed_cost / profit_multiplier /
|
||||
model_price / width / height / fps 字段,便于前端展示计费明细。
|
||||
"""
|
||||
w = max(1, int(width or 1))
|
||||
h = max(1, int(height or 1))
|
||||
effective_fps = int(fps or VIRAL_VIDEO_FPS)
|
||||
|
||||
prefix = _match_model_prefix(model)
|
||||
cfg = get_viral_video_model_config(prefix)
|
||||
res_key = _infer_resolution_key(w, h)
|
||||
billing = cfg.get("billing_mode", "token")
|
||||
key = (prefix, res_key, bool(has_video_input))
|
||||
price = VIRAL_VIDEO_MODEL_PRICES.get(key)
|
||||
if price is None:
|
||||
price = VIRAL_VIDEO_MODEL_PRICES.get(("seedance-2.5", res_key, False), 70.0)
|
||||
dur = max(1, int(duration_seconds or 15))
|
||||
if billing == "per_second":
|
||||
tokens = 0.0
|
||||
video_cost = dur * float(price)
|
||||
billing_unit = "second"
|
||||
else:
|
||||
if actual_tokens is not None and actual_tokens > 0:
|
||||
tokens = float(actual_tokens)
|
||||
else:
|
||||
tokens = dur * w * h * effective_fps / 1024.0
|
||||
video_cost = tokens / 1_000_000.0 * float(price)
|
||||
billing_unit = "token"
|
||||
|
||||
total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER
|
||||
credits = round(float(total), 2)
|
||||
breakdown = {
|
||||
"tokens": float(tokens),
|
||||
"video_cost": float(video_cost),
|
||||
"fixed_cost": float(VIRAL_VIDEO_FIXED_COST),
|
||||
"profit_multiplier": float(VIRAL_VIDEO_PROFIT_MULTIPLIER),
|
||||
"model_price": float(price),
|
||||
"model_key": prefix,
|
||||
"billing_mode": billing,
|
||||
"billing_unit": billing_unit,
|
||||
"width": int(w),
|
||||
"height": int(h),
|
||||
"fps": int(effective_fps),
|
||||
"duration": dur,
|
||||
}
|
||||
return credits, breakdown
|
||||
|
||||
|
||||
def calculate_viral_video_credits(
|
||||
duration_seconds: int,
|
||||
width: int,
|
||||
height: int,
|
||||
model: str = "seedance-2.5",
|
||||
has_video_input: bool = False,
|
||||
actual_tokens: int | None = None,
|
||||
fps: int = VIRAL_VIDEO_FPS,
|
||||
) -> float:
|
||||
"""计算爆款视频所需积分(1 积分 = 1 元),仅返回积分值(向后兼容包装器)。
|
||||
|
||||
内部调用 calculate_viral_video_credits_with_breakdown,仅返回 credits 部分,
|
||||
保持旧调用方签名与返回值类型不变。
|
||||
|
||||
公式:
|
||||
tokens = duration * width * height * fps / 1024
|
||||
video_cost = tokens / 1_000_000 * model_token_price
|
||||
total = round((video_cost + fixed_cost) * profit_multiplier, 2)
|
||||
若传入 actual_tokens 则用它替代计算值。
|
||||
"""
|
||||
credits, _ = calculate_viral_video_credits_with_breakdown(
|
||||
duration_seconds=duration_seconds,
|
||||
width=width,
|
||||
height=height,
|
||||
model=model,
|
||||
has_video_input=has_video_input,
|
||||
actual_tokens=actual_tokens,
|
||||
fps=fps,
|
||||
)
|
||||
return credits
|
||||
|
||||
|
||||
# ============ 场景定义 ============
|
||||
# 每个场景: base_points(基础积分), unit(计费单位), name(显示名称), dynamic(是否动态定价)
|
||||
# 说明:爆款视频(viral_video)走动态定价(预扣→结算多退少补),因此不使用 @points_gate
|
||||
# 装饰器,base_points=0,dynamic=True;前端展示场景列表时仍可看到。
|
||||
# 每个场景: base_points(基础积分), unit(计费单位), name(显示名称)
|
||||
|
||||
POINTS_SCENES: dict[str, dict] = {
|
||||
"ai_voice": {
|
||||
"base_points": 1,
|
||||
"unit": "分钟",
|
||||
"name": "AI 配音",
|
||||
"description": "AI 配音每分钟消耗 1 积分(免费用户上浮 15%,会员 8~9 折)",
|
||||
},
|
||||
"ai_video": {
|
||||
"base_points": 3,
|
||||
"unit": "条",
|
||||
"name": "智能混剪",
|
||||
"extra_per_30s": 1,
|
||||
"description": "智能混剪每条 3 积分起,视频超过 30 秒后每 30 秒加 1 积分;免费用户每日 2 条免费额度",
|
||||
},
|
||||
"ai_digital_human": {
|
||||
"base_points": 15,
|
||||
"unit": "分钟",
|
||||
"name": "AI 数字人",
|
||||
"description": "AI 数字人每分钟消耗 15 积分",
|
||||
},
|
||||
"voice_clone_train": {
|
||||
"base_points": 0,
|
||||
"unit": "次",
|
||||
@@ -347,16 +39,23 @@ POINTS_SCENES: dict[str, dict] = {
|
||||
"name": "声音克隆合成",
|
||||
"description": "克隆音色合成每分钟消耗 1 积分",
|
||||
},
|
||||
"viral_video": {
|
||||
"base_points": 0,
|
||||
"douyin_extract": {
|
||||
"base_points": 1,
|
||||
"unit": "次",
|
||||
"name": "爆款视频",
|
||||
"dynamic": True,
|
||||
"description": "爆款视频动态定价(按视频时长/分辨率/模型计算,预扣→结算多退少补)",
|
||||
"name": "抖音链接提取",
|
||||
"description": "抖音文案提取每次 1 积分",
|
||||
},
|
||||
"ai_rewrite": {"base_points": 1, "unit": "次", "name": "AI 改写文案", "description": "AI 改写文案每次 1 积分"},
|
||||
"ai_title": {
|
||||
"base_points": 1,
|
||||
"unit": "次",
|
||||
"name": "AI 标题生成",
|
||||
"description": "AI 生成标题每次 1 积分(免费用户实际上浮后 2 积分/次)",
|
||||
},
|
||||
"ai_cover": {"base_points": 1, "unit": "张", "name": "AI 封面生成", "description": "AI 封面生成每张 1 积分"},
|
||||
}
|
||||
|
||||
# 免费用户积分消耗上浮系数(仅对 voice_clone_synth 生效)
|
||||
# 免费用户积分消耗上浮系数
|
||||
FREE_USER_MULTIPLIER = 1.15
|
||||
|
||||
# ============ 积分包定义 ============
|
||||
@@ -382,6 +81,9 @@ MEMBER_DISCOUNT: dict[str, float] = {
|
||||
"yearly": 0.8,
|
||||
}
|
||||
|
||||
# 每日免费混剪次数(免费用户)
|
||||
DAILY_FREE_CLIP_LIMIT = 2
|
||||
|
||||
|
||||
def calculate_points_cost(
|
||||
scene_key: str,
|
||||
@@ -389,45 +91,48 @@ def calculate_points_cost(
|
||||
quantity: int = 1,
|
||||
duration_minutes: float = 0,
|
||||
member_type: str | None = None,
|
||||
) -> float:
|
||||
) -> int:
|
||||
"""计算指定场景的积分消耗。
|
||||
|
||||
Args:
|
||||
scene_key: 场景标识(当前支持 voice_clone_train/voice_clone_synth/viral_video;
|
||||
viral_video 为动态定价场景,此处返回 0,由业务侧调用
|
||||
calculate_viral_video_credits 手动计算)
|
||||
scene_key: 场景标识,如 "ai_voice"、"ai_video"
|
||||
is_member: 是否付费会员
|
||||
quantity: 数量(按次计费场景)
|
||||
duration_minutes: 时长分钟数(按时长计费场景)
|
||||
member_type: 会员类型 (monthly/quarterly/yearly),用于折扣
|
||||
|
||||
Returns:
|
||||
实际消耗积分(float;已含免费用户 ×1.15 上浮或会员折扣);免费/动态/已下线场景统一返回 0。
|
||||
实际消耗积分(已含免费用户 ×1.15 上浮或会员折扣)
|
||||
|
||||
Raises:
|
||||
ValueError: 未知场景标识
|
||||
"""
|
||||
scene = POINTS_SCENES.get(scene_key)
|
||||
if not scene:
|
||||
# 已下线/未注册的场景统一返回 0(免费),保持向后兼容
|
||||
return 0.0
|
||||
|
||||
# 动态定价场景(如 viral_video)由业务侧手动计算,这里统一返回 0
|
||||
if scene.get("dynamic"):
|
||||
return 0.0
|
||||
raise ValueError(f"Unknown points scene: {scene_key}")
|
||||
|
||||
base = scene["base_points"]
|
||||
if base == 0:
|
||||
return 0.0
|
||||
return 0
|
||||
|
||||
# —— 计算基础消耗 ——
|
||||
unit = scene["unit"]
|
||||
if unit == "分钟":
|
||||
total_base = base * max(1, math.ceil(duration_minutes))
|
||||
elif unit in ("次", "张"):
|
||||
elif unit in ("条", "次", "张"):
|
||||
total_base = base * quantity
|
||||
# 混剪特殊逻辑:视频超过 30s 后每 +30s 额外加 1 积分
|
||||
if scene_key == "ai_video" and duration_minutes > 0.5:
|
||||
extra_segments = math.ceil((duration_minutes * 60 - 30) / 30)
|
||||
if extra_segments > 0:
|
||||
total_base += scene.get("extra_per_30s", 1) * extra_segments
|
||||
else:
|
||||
total_base = base
|
||||
|
||||
# —— 会员折扣 / 免费用户上浮 ——
|
||||
if is_member and member_type and member_type in MEMBER_DISCOUNT:
|
||||
total_base = max(1, math.floor(total_base * MEMBER_DISCOUNT[member_type]))
|
||||
elif not is_member:
|
||||
total_base = math.ceil(total_base * FREE_USER_MULTIPLIER)
|
||||
|
||||
return float(total_base)
|
||||
return total_base
|
||||
|
||||
@@ -13,6 +13,7 @@ from typing import Any
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.domain.points_rules import (
|
||||
DAILY_FREE_CLIP_LIMIT,
|
||||
POINTS_PACKAGES,
|
||||
)
|
||||
|
||||
@@ -83,7 +84,7 @@ class PointsService:
|
||||
|
||||
# ──────────────── 余额检查 ────────────────
|
||||
|
||||
def check_balance(self, user_id: str, amount: float, db: Session) -> dict[str, Any]:
|
||||
def check_balance(self, user_id: str, amount: int, db: Session) -> dict[str, Any]:
|
||||
"""检查余额是否足够。"""
|
||||
account_data = self.get_or_create_account(user_id, db)
|
||||
balance = account_data["balance"]
|
||||
@@ -99,7 +100,7 @@ class PointsService:
|
||||
def deduct_points(
|
||||
self,
|
||||
user_id: str,
|
||||
amount: float,
|
||||
amount: int,
|
||||
source: str,
|
||||
db: Session,
|
||||
description: str = "",
|
||||
@@ -108,7 +109,7 @@ class PointsService:
|
||||
"""扣减积分(事务性:SELECT FOR UPDATE → 检查余额 → 扣减 → 流水 → 同步用户表)。
|
||||
|
||||
Returns:
|
||||
{"success": True/False, "balance": float, "transaction_id": str|None}
|
||||
{"success": True/False, "balance": int, "transaction_id": str|None}
|
||||
"""
|
||||
PointsAccountModel, PointsTransactionModel, _, _, UserModel = _get_models()
|
||||
|
||||
@@ -172,7 +173,7 @@ class PointsService:
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(
|
||||
"积分扣减失败: user_id=%s, amount=%.2f, source=%s",
|
||||
"积分扣减失败: user_id=%s, amount=%d, source=%s",
|
||||
user_id,
|
||||
amount,
|
||||
source,
|
||||
@@ -184,7 +185,7 @@ class PointsService:
|
||||
def add_points(
|
||||
self,
|
||||
user_id: str,
|
||||
amount: float,
|
||||
amount: int,
|
||||
source: str,
|
||||
db: Session,
|
||||
description: str = "",
|
||||
@@ -241,7 +242,7 @@ class PointsService:
|
||||
except Exception:
|
||||
db.rollback()
|
||||
logger.exception(
|
||||
"积分增加失败: user_id=%s, amount=%.2f, source=%s",
|
||||
"积分增加失败: user_id=%s, amount=%d, source=%s",
|
||||
user_id,
|
||||
amount,
|
||||
source,
|
||||
@@ -253,7 +254,7 @@ class PointsService:
|
||||
def refund_points(
|
||||
self,
|
||||
user_id: str,
|
||||
amount: float,
|
||||
amount: int,
|
||||
source: str,
|
||||
db: Session,
|
||||
ref_id: str = "",
|
||||
@@ -269,92 +270,6 @@ class PointsService:
|
||||
ref_id=ref_id,
|
||||
)
|
||||
|
||||
# ──────────────── 爆款视频(viral_video)动态定价 ────────────────
|
||||
|
||||
def deduct_viral_video(self, user_id: str, credits: float, job_id: str, db: Session) -> dict[str, Any]:
|
||||
"""爆款视频预扣积分(confirm-copy 阶段)。"""
|
||||
return self.deduct_points(
|
||||
user_id=user_id,
|
||||
amount=float(credits or 0),
|
||||
source="viral_video",
|
||||
db=db,
|
||||
description="爆款视频生成",
|
||||
ref_id=job_id,
|
||||
)
|
||||
|
||||
def settle_viral_video(
|
||||
self,
|
||||
user_id: str,
|
||||
estimated: float,
|
||||
actual: float,
|
||||
txn_id: str,
|
||||
db: Session,
|
||||
) -> dict[str, Any]:
|
||||
"""爆款视频完成后按实际 tokens 结算(多退少补)。
|
||||
|
||||
- actual < estimated: 退差额
|
||||
- actual > estimated: 补扣差额(余额不足时记 warning,不阻塞完成)
|
||||
- |diff| < 0.01: 不动
|
||||
"""
|
||||
diff = round(float(actual or 0) - float(estimated or 0), 2)
|
||||
if abs(diff) < 0.01:
|
||||
return {"success": True, "action": "none", "diff": 0.0}
|
||||
if diff < 0:
|
||||
refund = round(-diff, 2)
|
||||
try:
|
||||
res = self.refund_points(
|
||||
user_id=user_id,
|
||||
amount=refund,
|
||||
source="viral_video",
|
||||
db=db,
|
||||
ref_id=txn_id,
|
||||
description="爆款视频结算退费",
|
||||
)
|
||||
return {"success": bool(res.get("success")), "action": "refund", "diff": -refund, "amount": refund}
|
||||
except Exception:
|
||||
logger.exception("[viral_video] 结算退费异常 user_id=%s refund=%.2f", user_id, refund)
|
||||
return {"success": False, "action": "refund", "diff": -refund}
|
||||
else:
|
||||
extra = round(diff, 2)
|
||||
try:
|
||||
res = self.deduct_points(
|
||||
user_id=user_id,
|
||||
amount=extra,
|
||||
source="viral_video",
|
||||
db=db,
|
||||
description="爆款视频结算补扣",
|
||||
ref_id=txn_id,
|
||||
)
|
||||
if not res.get("success"):
|
||||
logger.warning(
|
||||
"[viral_video] 结算补扣余额不足 user_id=%s extra=%.2f balance=%s (不阻塞任务完成)",
|
||||
user_id,
|
||||
extra,
|
||||
res.get("balance"),
|
||||
)
|
||||
return {"success": bool(res.get("success")), "action": "deduct", "diff": extra, "amount": extra}
|
||||
except Exception:
|
||||
logger.exception("[viral_video] 结算补扣异常 user_id=%s extra=%.2f", user_id, extra)
|
||||
return {"success": False, "action": "deduct", "diff": extra}
|
||||
|
||||
def refund_viral_video(self, user_id: str, credits: float, txn_id: str, db: Session) -> dict[str, Any]:
|
||||
"""爆款视频失败全额退款。"""
|
||||
amount = float(credits or 0)
|
||||
if amount <= 0:
|
||||
return {"success": True, "action": "none", "amount": 0.0}
|
||||
try:
|
||||
return self.refund_points(
|
||||
user_id=user_id,
|
||||
amount=amount,
|
||||
source="viral_video",
|
||||
db=db,
|
||||
ref_id=txn_id,
|
||||
description="爆款视频失败退款",
|
||||
)
|
||||
except Exception:
|
||||
logger.exception("[viral_video] 失败退款异常 user_id=%s amount=%.2f", user_id, amount)
|
||||
return {"success": False, "action": "refund", "amount": amount}
|
||||
|
||||
# ──────────────── 流水查询 ────────────────
|
||||
|
||||
def get_transactions(
|
||||
@@ -409,16 +324,132 @@ class PointsService:
|
||||
"page_size": page_size,
|
||||
}
|
||||
|
||||
# ──────────────── 每日免费混剪额度(已下线:智能混剪全免费) ────────────────
|
||||
# ──────────────── 每日免费混剪额度 ────────────────
|
||||
|
||||
def _daily_key(self, user_id: str) -> str:
|
||||
"""生成 Redis 每日额度 key。格式: daily_usage:{user_id}:{YYYYMMDD}:free_clip"""
|
||||
today = datetime.now(UTC).strftime("%Y%m%d")
|
||||
return f"daily_usage:{user_id}:{today}:free_clip"
|
||||
|
||||
def check_daily_free_clip(self, user_id: str, db: Session) -> bool:
|
||||
"""检查今日是否还有免费混剪额度。
|
||||
|
||||
优先查 Redis,Redis 不可用时降级到 DB。
|
||||
"""
|
||||
redis_client = _get_redis_client()
|
||||
if redis_client:
|
||||
try:
|
||||
key = self._daily_key(user_id)
|
||||
current = redis_client.get(key)
|
||||
if current is None:
|
||||
return True
|
||||
return int(current) < DAILY_FREE_CLIP_LIMIT
|
||||
except Exception:
|
||||
logger.warning("Redis 不可用,降级到 DB 查询每日额度")
|
||||
|
||||
# 降级到 DB
|
||||
_, _, _, DailyUsageRecordModel, _ = _get_models()
|
||||
today_start = datetime.now(UTC).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
record = (
|
||||
db.query(DailyUsageRecordModel)
|
||||
.filter(
|
||||
DailyUsageRecordModel.user_id == user_id,
|
||||
DailyUsageRecordModel.usage_type == "free_clip",
|
||||
DailyUsageRecordModel.usage_date >= today_start,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if record is None:
|
||||
return True
|
||||
return record.count < DAILY_FREE_CLIP_LIMIT
|
||||
|
||||
def record_daily_free_clip(self, user_id: str, db: Session) -> bool:
|
||||
"""记录使用一次免费混剪。
|
||||
|
||||
先 INCR Redis;如果超限回退 Redis。DB 使用 upsert 语义(唯一约束)。
|
||||
"""
|
||||
redis_client = _get_redis_client()
|
||||
if redis_client:
|
||||
try:
|
||||
key = self._daily_key(user_id)
|
||||
new_count = redis_client.incr(key)
|
||||
if new_count == 1:
|
||||
redis_client.expire(key, 48 * 3600) # TTL 48h
|
||||
if new_count <= DAILY_FREE_CLIP_LIMIT:
|
||||
return True
|
||||
# 超限,回退 Redis
|
||||
redis_client.decr(key)
|
||||
except Exception:
|
||||
logger.warning("Redis 不可用,降级到 DB 记录每日额度")
|
||||
|
||||
# 降级/兜底到 DB(upsert 语义)
|
||||
_, _, _, DailyUsageRecordModel, _ = _get_models()
|
||||
today_start = datetime.now(UTC).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
|
||||
record = (
|
||||
db.query(DailyUsageRecordModel)
|
||||
.filter(
|
||||
DailyUsageRecordModel.user_id == user_id,
|
||||
DailyUsageRecordModel.usage_type == "free_clip",
|
||||
DailyUsageRecordModel.usage_date >= today_start,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
if record is None:
|
||||
if DAILY_FREE_CLIP_LIMIT <= 0:
|
||||
return False
|
||||
record = DailyUsageRecordModel(
|
||||
id=uuid.uuid4().hex,
|
||||
user_id=user_id,
|
||||
usage_type="free_clip",
|
||||
usage_date=datetime.now(UTC),
|
||||
count=1,
|
||||
)
|
||||
db.add(record)
|
||||
else:
|
||||
if record.count >= DAILY_FREE_CLIP_LIMIT:
|
||||
return False
|
||||
record.count += 1
|
||||
|
||||
db.commit()
|
||||
return True
|
||||
|
||||
def get_daily_usage(self, user_id: str, db: Session) -> dict[str, Any]:
|
||||
"""查询今日免费额度使用情况(智能混剪已全免费,返回 unlimited)。"""
|
||||
"""查询今日免费额度使用情况。"""
|
||||
redis_client = _get_redis_client()
|
||||
used = 0
|
||||
|
||||
if redis_client:
|
||||
try:
|
||||
key = self._daily_key(user_id)
|
||||
val = redis_client.get(key)
|
||||
used = int(val) if val else 0
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if used == 0:
|
||||
# 从 DB 查
|
||||
_, _, _, DailyUsageRecordModel, _ = _get_models()
|
||||
today_start = datetime.now(UTC).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
record = (
|
||||
db.query(DailyUsageRecordModel)
|
||||
.filter(
|
||||
DailyUsageRecordModel.user_id == user_id,
|
||||
DailyUsageRecordModel.usage_type == "free_clip",
|
||||
DailyUsageRecordModel.usage_date >= today_start,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
used = record.count if record else 0
|
||||
|
||||
now = datetime.now(UTC)
|
||||
tomorrow = (now + timedelta(days=1)).replace(hour=0, minute=0, second=0, microsecond=0)
|
||||
|
||||
return {
|
||||
"free_clips_used": 0,
|
||||
"free_clips_limit": -1, # -1 表示 unlimited
|
||||
"free_clips_remaining": -1,
|
||||
"free_clips_used": used,
|
||||
"free_clips_limit": DAILY_FREE_CLIP_LIMIT,
|
||||
"free_clips_remaining": max(0, DAILY_FREE_CLIP_LIMIT - used),
|
||||
"reset_at": tomorrow.isoformat(),
|
||||
}
|
||||
|
||||
|
||||
@@ -69,6 +69,8 @@ class PromptType(StrEnum):
|
||||
STYLE_CONSTRAINT = "style_constraint"
|
||||
|
||||
|
||||
CREDITS_VIRAL_VIDEO_COST = 50
|
||||
|
||||
STAGE_LABELS = {
|
||||
ViralVideoStage.IMAGE_ANALYSIS: "图片分析",
|
||||
ViralVideoStage.VIDEO_ANALYSIS: "视频风格分析",
|
||||
@@ -119,10 +121,7 @@ class ViralVideoJob:
|
||||
phase_message: str = "" # 阶段中文提示文案,前端轮询直接展示
|
||||
heartbeat_at: datetime | None = None # worker 心跳时间,用于超时僵尸任务检测
|
||||
result_video_url: str = ""
|
||||
video_resolution: str = "720p"
|
||||
credits_prepaid: float = 0.0
|
||||
credits_transaction_id: str = ""
|
||||
credits_cost: float = 0.0
|
||||
credits_cost: int = 0
|
||||
error_msg: str = ""
|
||||
retry_count: int = 0
|
||||
started_at: datetime | None = None
|
||||
|
||||
@@ -185,6 +185,19 @@ def _execute_with_gate_impl(
|
||||
is_member = getattr(user, "is_member", False)
|
||||
member_type = getattr(user, "member_type", None)
|
||||
|
||||
if scene_key == "ai_video":
|
||||
from packages.domain.points_service import PointsService
|
||||
|
||||
svc = PointsService()
|
||||
if not is_member:
|
||||
if svc.check_daily_free_clip(user.id, db):
|
||||
svc.record_daily_free_clip(user.id, db)
|
||||
kwargs["_points_deducted"] = 0
|
||||
kwargs["_is_free_quota"] = True
|
||||
if is_async:
|
||||
return _run_async_impl(func, args, _filter_kwargs_impl(func, kwargs))
|
||||
return func(*args, **_filter_kwargs_impl(func, kwargs))
|
||||
|
||||
if per_unit is not None:
|
||||
total_points = per_unit
|
||||
else:
|
||||
|
||||
+26
-466
@@ -22,153 +22,9 @@ import httpx
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
# 网络/超时类异常父类集合:覆盖 Timeout/Connect/Network/ReadTimeout/WriteTimeout/PoolTimeout
|
||||
_HTTP_NETWORK_ERRORS = ()
|
||||
try:
|
||||
_HTTP_NETWORK_ERRORS = (httpx.TimeoutException, httpx.NetworkError)
|
||||
except Exception:
|
||||
_HTTP_NETWORK_ERRORS = (Exception,)
|
||||
|
||||
_HTTP_STATUS_ERROR = httpx.HTTPStatusError if hasattr(httpx, "HTTPStatusError") else Exception
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# 视频模型 ID 解析逻辑(#2159 多模型支持,#2169 接入即梦)。
|
||||
# 内部使用简短别名(seedance-2.5 / seedance-2.0 / wan-3.0 / jimeng-3.0 等)做 PRICING key;
|
||||
# 实际调用时按 VIRAL_VIDEO_MODEL_CONFIG(domain/points_rules.py)的 provider/model_id 分发。
|
||||
# - provider=doubao → 火山方舟 Seedance
|
||||
# - provider=dashscope → 阿里云 DashScope(Wan 系列)
|
||||
# - provider=jimeng → 火山引擎即梦 cvtob(jimeng_i2v_first_v30,真人参考图走 C 端审核)
|
||||
|
||||
|
||||
def _resolve_video_provider_and_id(model: str | None) -> tuple[str, str, dict]:
|
||||
"""把内部 model key 解析成 (provider, model_id, cfg)。
|
||||
|
||||
- provider: "doubao" | "dashscope" | "jimeng"
|
||||
- model_id: 对应 API 的真实模型 ID
|
||||
- cfg: VIRAL_VIDEO_MODEL_CONFIG 条目
|
||||
未识别或空值回落到默认 seedance-2.5。已经是 doubao-/ep- 开头的完整 ID 视为 doubao provider;
|
||||
"jimeng" 开头视为 jimeng provider(内部兜底,不暴露给前端)。
|
||||
"""
|
||||
from packages.domain.points_rules import get_viral_video_model_config
|
||||
|
||||
settings = get_shared_settings()
|
||||
default_id = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
|
||||
m = (model or "").strip()
|
||||
if not m:
|
||||
cfg = get_viral_video_model_config("seedance-2.5")
|
||||
return "doubao", default_id, cfg
|
||||
# 已经是 doubao-/ep- 开头:直接透传,默认视为 doubao provider
|
||||
if m.startswith("doubao-") or m.startswith("ep-"):
|
||||
return "doubao", m, {"provider": "doubao", "model_id": m, "supports_audio": True}
|
||||
# 显式 jimeng 关键字:路由到即梦(内部兜底通道使用)
|
||||
if m.startswith("jimeng"):
|
||||
cfg = get_viral_video_model_config("jimeng-3.0")
|
||||
return "jimeng", cfg.get("model_id", "jimeng_i2v_first_v30"), cfg
|
||||
# 别名 → 从 domain config 查
|
||||
cfg = get_viral_video_model_config(m)
|
||||
provider = cfg.get("provider", "doubao")
|
||||
resolved_id = cfg.get("model_id", "")
|
||||
if not resolved_id:
|
||||
logger.warning("[ai_client] model %r 无 model_id,回落到默认 %s", m, default_id)
|
||||
return "doubao", default_id, cfg
|
||||
return provider, resolved_id, cfg
|
||||
|
||||
|
||||
def _resolve_video_model_id(model: str | None) -> str:
|
||||
"""兼容旧调用:只返回 doubao model_id。wan/dashscope 调用方应直接用 _resolve_video_provider_and_id。"""
|
||||
_provider, mid, _cfg = _resolve_video_provider_and_id(model)
|
||||
return mid
|
||||
|
||||
|
||||
# ── 视频错误分类(给前端/用户展示友好提示)────────────────────────────
|
||||
|
||||
|
||||
def _classify_video_error(status_code: int, body: str, err: Exception | None) -> tuple[str, str]:
|
||||
"""根据 HTTP 状态码和响应 body 判断错误类型。
|
||||
|
||||
返回 (error_code, user_message):
|
||||
- error_code: 机器可读的错误码("portrait_intercept" / "quota_exceeded" / "model_not_found"
|
||||
/ "invalid_param" / "auth_error" / "rate_limit" / "network_error" / "task_failed" / "unknown")
|
||||
- user_message: 给用户看的中文提示
|
||||
"""
|
||||
body_lower = (body or "").lower()
|
||||
code_in_body = ""
|
||||
msg_in_body = ""
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
parsed = _json.loads(body or "{}")
|
||||
if isinstance(parsed, dict):
|
||||
err_obj = parsed.get("error") or {}
|
||||
if isinstance(err_obj, dict):
|
||||
code_in_body = str(err_obj.get("code", "") or "")
|
||||
msg_in_body = str(err_obj.get("message", "") or err_obj.get("msg", "") or "")
|
||||
else:
|
||||
msg_in_body = str(parsed.get("message", "") or "")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# 真人肖像/内容安全拦截
|
||||
if (
|
||||
status_code == 400
|
||||
and any(
|
||||
kw in body_lower
|
||||
for kw in ("portrait", "real_face", "human_face", "真人", "肖像", "人脸", "privacy", "real person", "face")
|
||||
)
|
||||
) or (
|
||||
"content" in body_lower
|
||||
and ("risk" in body_lower or "block" in body_lower or "reject" in body_lower)
|
||||
and status_code == 400
|
||||
):
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"参考素材包含真人照片被安全策略拦截,AI视频模型暂不支持上传真人照片作为参考图,请移除真人图片后重试。",
|
||||
)
|
||||
|
||||
# 配额/计费问题
|
||||
if status_code in (402, 429) or any(
|
||||
kw in body_lower for kw in ("quota", "billing", "insufficient", "欠费", "余额", "限流", "rate limit")
|
||||
):
|
||||
if "rate" in body_lower or status_code == 429:
|
||||
return "rate_limit", "视频生成服务当前繁忙(限流),请稍等1-2分钟后重试。"
|
||||
return "quota_exceeded", "视频生成服务配额不足,请联系管理员充值或稍后重试。"
|
||||
|
||||
# 模型/Endpoint 不存在
|
||||
if status_code == 404 or any(
|
||||
kw in body_lower for kw in ("model not found", "endpoint not found", "不存在", "not found", "model_not_exist")
|
||||
):
|
||||
return "model_not_found", f"视频模型未开通或模型ID无效({code_in_body or ''}),请联系管理员。"
|
||||
|
||||
# 鉴权失败
|
||||
if status_code in (401, 403):
|
||||
return "auth_error", "视频生成服务鉴权失败(API Key无效或过期),请联系管理员。"
|
||||
|
||||
# 任务本身失败(轮询阶段拿到 status=failed)
|
||||
if err and "task failed" in str(err).lower():
|
||||
detail = msg_in_body or str(err)[:200]
|
||||
# 失败原因里再细分真人拦截
|
||||
if any(kw in detail.lower() for kw in ("portrait", "真人", "肖像", "人脸", "content_risk")):
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"视频内容被安全策略拦截(疑似包含真人肖像),请更换参考图或调整文案后重试。",
|
||||
)
|
||||
return "task_failed", f"视频生成失败:{detail}"
|
||||
|
||||
# 参数错误
|
||||
if status_code == 400:
|
||||
return "invalid_param", f"视频生成参数错误:{msg_in_body or body[:200]}"
|
||||
|
||||
# 网络/连接问题
|
||||
if status_code == 0:
|
||||
return "network_error", "视频生成服务连接失败(网络超时),请稍后重试。"
|
||||
|
||||
# 默认
|
||||
detail = msg_in_body or (str(err) if err else "") or body[:200]
|
||||
return "unknown", f"视频生成失败(HTTP {status_code}):{detail}"
|
||||
|
||||
|
||||
class DoubaoClient:
|
||||
"""豆包大模型 API 客户端.
|
||||
|
||||
@@ -186,9 +42,6 @@ class DoubaoClient:
|
||||
self.vision_model: str = settings.doubao_vision_model
|
||||
self.vision_lite_model: str = settings.doubao_vision_lite_model
|
||||
self.fast_model: str = settings.doubao_fast_model
|
||||
self.embedding_model: str = settings.doubao_embedding_model
|
||||
# 最近一次视频生成的详细错误(error_code + user_message + raw detail),供上层读取后展示给用户
|
||||
self.last_video_error: dict = {}
|
||||
|
||||
def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None:
|
||||
"""调用豆包文本 Embedding API,返回浮点向量;失败返回 None。"""
|
||||
@@ -201,7 +54,7 @@ class DoubaoClient:
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload: dict[str, Any] = {
|
||||
"model": self.embedding_model,
|
||||
"model": getattr(self, "embedding_model", None) or "doubao-embedding-large-text-240915",
|
||||
"input": text.strip(),
|
||||
"encoding_format": "float",
|
||||
}
|
||||
@@ -408,12 +261,8 @@ class DoubaoClient:
|
||||
reference_images: list[str] | None = None,
|
||||
reference_audios: list[str] | None = None,
|
||||
reference_videos: list[str] | None = None,
|
||||
) -> dict | None:
|
||||
"""调用 Seedance 2.5 生视频(异步任务→轮询→下载)。
|
||||
|
||||
成功返回 {"video_path": str, "usage": dict | None},失败返回 None。
|
||||
失败时把详细错误信息(HTTP状态码、响应 body、分类后的用户提示)写入 self.last_video_error,
|
||||
上层可通过 get_last_video_error() 读取并展示给用户,不再笼统显示"返回为空"。
|
||||
) -> str | None:
|
||||
"""调用 Seedance 2.5 生视频(异步任务→轮询→下载),返回本地 MP4 路径;失败返回 None。
|
||||
|
||||
【v1.6.1 修复】严格按官方 content 数组协议构造请求:
|
||||
- 所有参考(图/音/视)必须放进 content 数组并带 role 字段,不能放顶层 reference_audios/reference_videos(非官方字段,会被忽略或导致异常)。
|
||||
@@ -421,83 +270,17 @@ class DoubaoClient:
|
||||
判定:传了参考音频/视频或 ≥1 张多参考图时,走 omni_reference(首张图 role=reference_image);纯首帧无参考时走 first_frame(ratio 强制 adaptive)。
|
||||
- 创建任务若因 ratio 报错(HTTP 400),自动回退到 ratio=adaptive 重试一次。
|
||||
"""
|
||||
# 每次调用前清空上次错误
|
||||
self.last_video_error = {}
|
||||
|
||||
if not self.is_available:
|
||||
self.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "视频生成服务未配置(API Key 缺失),请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": "DoubaoClient not available (api_key empty)",
|
||||
}
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
self.last_video_error = {
|
||||
"error_code": "invalid_param",
|
||||
"user_message": "视频生成提示词不能为空。",
|
||||
"status_code": 0,
|
||||
"detail": "empty prompt",
|
||||
}
|
||||
return None
|
||||
|
||||
settings = get_shared_settings()
|
||||
poll_interval = getattr(settings, "doubao_video_poll_interval", 10) or 10
|
||||
# 收紧总超时:轮询 8min + 下载 2min = 最长 ~10min,防止出现 20min 卡死
|
||||
total_timeout = getattr(settings, "doubao_video_timeout", 480) or 480
|
||||
# 内部 key → (provider, 实际模型 ID, cfg),按 provider 分发
|
||||
provider, video_model, model_cfg = _resolve_video_provider_and_id(model)
|
||||
if provider == "dashscope":
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
ds = get_dashscope_client()
|
||||
if ds is None:
|
||||
err_msg = "DashScope client 不可用(未配置 DASHSCOPE_API_KEY)"
|
||||
logger.error("%s, video_model=%s", err_msg, model)
|
||||
self.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "Wan 3.0 视频模型未配置 API Key,请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": err_msg,
|
||||
}
|
||||
return None
|
||||
try:
|
||||
# DashScope 客户端也设置 last_video_error 语义(如果它支持)
|
||||
if hasattr(ds, "last_video_error"):
|
||||
ds.last_video_error = {}
|
||||
result = ds.video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
output_dir=output_dir,
|
||||
model=video_model,
|
||||
)
|
||||
if not result and hasattr(ds, "last_video_error") and ds.last_video_error:
|
||||
self.last_video_error = dict(ds.last_video_error)
|
||||
return result
|
||||
except Exception as de:
|
||||
logger.error("DashScope video_generation 异常: %s", de, exc_info=True)
|
||||
self.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": f"Wan 3.0 视频生成异常:{de!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(de),
|
||||
}
|
||||
return None
|
||||
|
||||
if provider == "jimeng":
|
||||
# #2169: 即梦 cvtob(jimeng_i2v_first_v30)— 真人参考图兜底通道
|
||||
return self._call_jimeng_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
output_dir=output_dir,
|
||||
generate_audio=generate_audio,
|
||||
)
|
||||
default_video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
|
||||
video_model = model or default_video_model
|
||||
|
||||
ref_audios = [u for u in (reference_audios or [])[:10] if u and isinstance(u, str)]
|
||||
ref_videos = [u for u in (reference_videos or [])[:3] if u and isinstance(u, str)]
|
||||
@@ -522,7 +305,7 @@ class DoubaoClient:
|
||||
}
|
||||
)
|
||||
else:
|
||||
# 纯首帧:显式 role=first_frame
|
||||
# 纯首帧:不带 role,服务端识别为 first_frame(或显式 role=first_frame)
|
||||
content.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
@@ -563,38 +346,22 @@ class DoubaoClient:
|
||||
len(ref_audios),
|
||||
len(ref_videos),
|
||||
)
|
||||
# 打印完整 payload 便于排查(截断 prompt)
|
||||
debug_payload = dict(create_payload)
|
||||
if "content" in debug_payload:
|
||||
dbg_content = []
|
||||
for item in debug_payload["content"]:
|
||||
item_copy = dict(item)
|
||||
if item_copy.get("type") == "text" and isinstance(item_copy.get("text"), str):
|
||||
item_copy["text"] = item_copy["text"][:200] + ("..." if len(item_copy["text"]) > 200 else "")
|
||||
dbg_content.append(item_copy)
|
||||
debug_payload["content"] = dbg_content
|
||||
logger.info("Seedance 创建任务 payload: %s", json_safe_dumps(debug_payload))
|
||||
|
||||
def _do_create(payload: dict) -> tuple[str | None, Exception | None, int, str]:
|
||||
"""返回 (task_id, last_err, status_code, body_text)。"""
|
||||
last_err: Exception | None = None
|
||||
last_sc = 0
|
||||
last_body = ""
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
resp = httpx.post(create_url, headers=headers, json=payload, timeout=self.timeout)
|
||||
sc = int(getattr(resp, "status_code", 0) or 0)
|
||||
body = (getattr(resp, "text", "") or "")[:2000]
|
||||
last_sc = sc
|
||||
last_body = body
|
||||
body = (getattr(resp, "text", "") or "")[:1500]
|
||||
if sc >= 400:
|
||||
logger.error("Seedance 创建任务 HTTP %d: body=%s", sc, body)
|
||||
try:
|
||||
resp.raise_for_status()
|
||||
except Exception as ee:
|
||||
last_err = ee
|
||||
if attempt < self.max_retries and sc >= 500:
|
||||
# 仅 5xx 重试,4xx 不重试(参数/鉴权/配额错误重试无意义)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
return None, last_err, sc, body
|
||||
@@ -603,19 +370,9 @@ class DoubaoClient:
|
||||
if tid:
|
||||
return tid, None, sc, body
|
||||
last_err = RuntimeError(f"create ok but no id: {str(data)[:300]}")
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
last_err = ne
|
||||
last_sc = 0
|
||||
last_body = f"network error: {ne}"
|
||||
logger.warning(
|
||||
"Seedance 创建网络异常(%s),重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
except Exception as e:
|
||||
last_err = e
|
||||
if attempt < self.max_retries and not isinstance(e, _HTTP_STATUS_ERROR):
|
||||
if attempt < self.max_retries:
|
||||
wait = 0.5 * (2**attempt)
|
||||
logger.warning(
|
||||
"Seedance 创建任务失败,%.1fs 后重试 (%d/%d): %s",
|
||||
@@ -625,7 +382,7 @@ class DoubaoClient:
|
||||
e,
|
||||
)
|
||||
time.sleep(wait)
|
||||
return None, last_err, last_sc, last_body
|
||||
return None, last_err, 0, ""
|
||||
|
||||
# 第一次尝试
|
||||
task_id, last_err, sc, body = _do_create(create_payload)
|
||||
@@ -644,53 +401,17 @@ class DoubaoClient:
|
||||
logger.warning("Seedance 创建因 ratio 失败,回退 ratio=adaptive 重试")
|
||||
create_payload["ratio"] = "adaptive"
|
||||
task_id, last_err, sc2, body2 = _do_create(create_payload)
|
||||
if task_id:
|
||||
sc, body = sc2, body2
|
||||
else:
|
||||
# 保留第二次的错误信息
|
||||
sc, body = sc2, body2
|
||||
|
||||
if not task_id:
|
||||
err_code, user_msg = _classify_video_error(sc, body, last_err)
|
||||
self.last_video_error = {
|
||||
"error_code": err_code,
|
||||
"user_message": user_msg,
|
||||
"status_code": sc,
|
||||
"detail": (body or "")[:500] or (str(last_err) if last_err else ""),
|
||||
"model": video_model,
|
||||
"base_url": self.base_url,
|
||||
}
|
||||
logger.error(
|
||||
"Seedance 创建任务最终失败: model=%s base_url=%s status=%d code=%s err=%s body=%s",
|
||||
"Seedance 创建任务最终失败: model=%s base_url=%s err=%s body=%s 【排查】"
|
||||
"1) 方舟控制台已开通 doubao-seedance-2-5-260628;2) API Key 有该模型权限;"
|
||||
"3) DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3;4) 参考素材 URL 公网可访问。",
|
||||
video_model,
|
||||
self.base_url,
|
||||
sc,
|
||||
err_code,
|
||||
last_err,
|
||||
(body or "")[:500],
|
||||
)
|
||||
# #2169: 方舟返回 portrait_intercept 且有参考图 → 自动切即梦重试一次(保留首帧图)
|
||||
if err_code == "portrait_intercept" and image_url:
|
||||
logger.warning(
|
||||
"[viral-video] Seedance 真人拦截(code=%s),自动切即梦通道重试(首帧图) img=%s",
|
||||
err_code,
|
||||
bool(image_url),
|
||||
)
|
||||
jm_result = self._call_jimeng_video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
output_dir=output_dir,
|
||||
generate_audio=False, # 即梦 i2v 不带音频,音频由后续 ffmpeg 合成
|
||||
_portrait_fallback=True,
|
||||
)
|
||||
if jm_result is not None:
|
||||
return jm_result
|
||||
# 即梦也失败了,保留即梦的 last_video_error(已经由 _call_jimeng 设置)
|
||||
logger.error("[viral-video] 即梦通道重试也失败: %s", self.last_video_error)
|
||||
return None
|
||||
return None
|
||||
|
||||
logger.info("Seedance 任务已创建: task_id=%s ratio=%s", task_id, create_payload["ratio"])
|
||||
@@ -699,50 +420,33 @@ class DoubaoClient:
|
||||
poll_url = f"{create_url}/{task_id}"
|
||||
deadline = time.time() + total_timeout
|
||||
video_url: str | None = None
|
||||
usage: dict | None = None
|
||||
last_status: str = "queued"
|
||||
poll_count = 0
|
||||
last_poll_body: str = ""
|
||||
last_poll_sc: int = 0
|
||||
while time.time() < deadline:
|
||||
poll_count += 1
|
||||
try:
|
||||
resp = httpx.get(poll_url, headers=headers, timeout=self.timeout)
|
||||
last_poll_sc = int(getattr(resp, "status_code", 200) or 200)
|
||||
last_poll_body = (getattr(resp, "text", "") or "")[:1500]
|
||||
if last_poll_sc >= 400:
|
||||
logger.warning("Seedance 轮询 HTTP %d: %s", last_poll_sc, last_poll_body[:300])
|
||||
if poll_count < 3:
|
||||
time.sleep(poll_interval)
|
||||
continue
|
||||
last_err = RuntimeError(f"poll HTTP {last_poll_sc}: {last_poll_body[:200]}")
|
||||
break
|
||||
try:
|
||||
if int(getattr(resp, "status_code", 200)) >= 400:
|
||||
resp.raise_for_status()
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
data = resp.json()
|
||||
status = data.get("status", "")
|
||||
last_status = status
|
||||
if status == "succeeded":
|
||||
content_obj = data.get("content") or {}
|
||||
video_url = content_obj.get("video_url")
|
||||
usage = data.get("usage") or content_obj.get("usage") or None
|
||||
if video_url:
|
||||
logger.info("Seedance 任务成功: task_id=%s polls=%d usage=%s", task_id, poll_count, usage)
|
||||
logger.info("Seedance 任务成功: task_id=%s polls=%d", task_id, poll_count)
|
||||
break
|
||||
# 成功但没 video_url:记录完整响应便于排查
|
||||
logger.error(
|
||||
"Seedance succeeded 但无 video_url: task_id=%s full_response=%s",
|
||||
task_id,
|
||||
str(data)[:1000],
|
||||
)
|
||||
last_err = RuntimeError("task succeeded but no video_url in response")
|
||||
last_poll_body = str(data)[:1000]
|
||||
last_err = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}")
|
||||
logger.error("Seedance succeeded 但无 video_url: %s", last_err)
|
||||
break
|
||||
if status == "failed":
|
||||
err = data.get("error") or {}
|
||||
err_code = str(err.get("code", "") or "")
|
||||
err_msg = str(err.get("message", "") or err.get("msg", "") or "")
|
||||
last_err = RuntimeError(f"task failed: code={err_code} msg={err_msg}")
|
||||
logger.error("Seedance 任务失败 task_id=%s code=%s msg=%s", task_id, err_code, err_msg)
|
||||
last_poll_body = str(data)[:1000]
|
||||
last_err = RuntimeError(f"task failed: code={err.get('code','')} msg={err.get('message','')}")
|
||||
logger.error("Seedance 任务失败 task_id=%s: %s", task_id, last_err)
|
||||
break
|
||||
if status in ("expired", "cancelled"):
|
||||
last_err = RuntimeError(f"task {status}")
|
||||
@@ -753,39 +457,18 @@ class DoubaoClient:
|
||||
logger.info("Seedance 轮询中: task_id=%s status=%s polls=%d", task_id, status, poll_count)
|
||||
except httpx.HTTPStatusError as e:
|
||||
last_err = e
|
||||
last_poll_sc = e.response.status_code
|
||||
last_poll_body = (e.response.text or "")[:500]
|
||||
logger.warning("Seedance 轮询 HTTP %d: %s", e.response.status_code, last_poll_body[:300])
|
||||
logger.warning("Seedance 轮询 HTTP %d: %s", e.response.status_code, (e.response.text or "")[:300])
|
||||
except Exception as e:
|
||||
last_err = e
|
||||
logger.debug("Seedance 轮询异常: %s", e)
|
||||
time.sleep(poll_interval)
|
||||
|
||||
if not video_url:
|
||||
# 区分轮询超时 vs 任务失败
|
||||
if last_status in ("queued", "running", "pending") and poll_count > 0 and time.time() >= deadline:
|
||||
err_code, user_msg = (
|
||||
"network_error",
|
||||
f"视频生成超时(>{total_timeout}s),任务仍在排队,请稍后重试或联系管理员。",
|
||||
)
|
||||
detail = f"timeout after {total_timeout}s, polls={poll_count}, last_status={last_status}"
|
||||
else:
|
||||
err_code, user_msg = _classify_video_error(last_poll_sc, last_poll_body, last_err)
|
||||
detail = (last_poll_body or "")[:500] or (str(last_err) if last_err else f"last_status={last_status}")
|
||||
self.last_video_error = {
|
||||
"error_code": err_code,
|
||||
"user_message": user_msg,
|
||||
"status_code": last_poll_sc,
|
||||
"detail": detail,
|
||||
"task_id": task_id,
|
||||
"last_status": last_status,
|
||||
}
|
||||
logger.error(
|
||||
"Seedance 任务未成功: task_id=%s last_status=%s polls=%d code=%s err=%s (总等待 %.0fs)",
|
||||
"Seedance 任务未成功: task_id=%s last_status=%s polls=%d err=%s (总等待 %.0fs)",
|
||||
task_id,
|
||||
last_status,
|
||||
poll_count,
|
||||
err_code,
|
||||
last_err,
|
||||
total_timeout,
|
||||
)
|
||||
@@ -816,135 +499,12 @@ class DoubaoClient:
|
||||
os.remove(local_path)
|
||||
except Exception:
|
||||
pass
|
||||
self.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": "视频生成成功但下载文件为空,请稍后重试。",
|
||||
"status_code": 0,
|
||||
"detail": f"downloaded 0 bytes from {video_url[:120]}",
|
||||
}
|
||||
return None
|
||||
return {"video_path": local_path, "usage": usage}
|
||||
return local_path
|
||||
except Exception as e:
|
||||
logger.error("Seedance 视频下载失败: %s", e, exc_info=True)
|
||||
self.last_video_error = {
|
||||
"error_code": "network_error",
|
||||
"user_message": f"视频下载失败:{e!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(e),
|
||||
}
|
||||
return None
|
||||
|
||||
def _call_jimeng_video_generation(
|
||||
self,
|
||||
*,
|
||||
prompt: str,
|
||||
image_url: str | None,
|
||||
duration: int,
|
||||
ratio: str | None,
|
||||
resolution: str,
|
||||
output_dir: str | None,
|
||||
generate_audio: bool = False,
|
||||
_portrait_fallback: bool = False,
|
||||
) -> dict | None:
|
||||
"""#2169: 调用即梦 cvtob 客户端做图生视频(真人参考图兜底通道)。
|
||||
|
||||
- 即梦 i2v 首帧接口只接受 1 张图、无原生音频(返回无声视频,音频由 ffmpeg 后合)。
|
||||
- 成功返回 {"video_path": str, "usage": {...}};失败写 self.last_video_error 并返回 None。
|
||||
- _portrait_fallback=True 时在日志里标注是从方舟拦截切过来的。
|
||||
"""
|
||||
from packages.shared.jimeng_client import get_jimeng_client
|
||||
|
||||
jm = get_jimeng_client()
|
||||
if jm is None:
|
||||
detail = "即梦 client 不可用(JIMENG_AK/SK 未配置)"
|
||||
if _portrait_fallback:
|
||||
# 从真人拦截切过来但即梦没配,仍把错误归到 portrait_intercept,让上层提示用户
|
||||
self.last_video_error = {
|
||||
"error_code": "portrait_intercept",
|
||||
"user_message": "参考素材包含真人照片被安全策略拦截,即梦兜底通道未启用,请联系管理员配置 JIMENG_AK/SK。",
|
||||
"status_code": 0,
|
||||
"detail": detail,
|
||||
"provider": "jimeng",
|
||||
}
|
||||
else:
|
||||
self.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "即梦视频通道未配置(JIMENG_AK/SK 缺失),请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": detail,
|
||||
"provider": "jimeng",
|
||||
}
|
||||
logger.error("[jimeng] %s, portrait_fallback=%s", detail, _portrait_fallback)
|
||||
return None
|
||||
if not image_url:
|
||||
self.last_video_error = {
|
||||
"error_code": "invalid_param",
|
||||
"user_message": "即梦图生视频必须提供参考图片。",
|
||||
"status_code": 0,
|
||||
"detail": "empty image_url for jimeng i2v",
|
||||
"provider": "jimeng",
|
||||
}
|
||||
return None
|
||||
# 即梦 i2v 无声视频,generate_audio 强制 False
|
||||
jm.last_video_error = {}
|
||||
tag = "[portrait-fallback→jimeng]" if _portrait_fallback else "[jimeng-direct]"
|
||||
logger.info("%s 调用即梦: dur=%s ratio=%s res=%s img=%s", tag, duration, ratio, resolution, bool(image_url))
|
||||
try:
|
||||
result = jm.video_generation(
|
||||
prompt=prompt,
|
||||
image_url=image_url,
|
||||
duration=duration,
|
||||
ratio=ratio,
|
||||
resolution=resolution,
|
||||
output_dir=output_dir,
|
||||
generate_audio=False,
|
||||
)
|
||||
except Exception as je:
|
||||
logger.error("%s 即梦 video_generation 异常: %s", tag, je, exc_info=True)
|
||||
self.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": f"即梦视频生成异常:{je!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(je),
|
||||
"provider": "jimeng",
|
||||
}
|
||||
return None
|
||||
if not result and jm.last_video_error:
|
||||
# 透传即梦错误;如果即梦也返回 portrait_intercept,说明图片真的有问题,直接给用户
|
||||
jm_err = dict(jm.last_video_error)
|
||||
jm_err["provider"] = "jimeng"
|
||||
if _portrait_fallback and jm_err.get("error_code") == "portrait_intercept":
|
||||
jm_err["user_message"] = (
|
||||
"参考素材真人肖像审核未通过(方舟+即梦双通道均被拦截),请更换非真人或授权清晰的照片后重试。"
|
||||
)
|
||||
self.last_video_error = jm_err
|
||||
return None
|
||||
if result:
|
||||
# 补充 usage 里的 provider 标记
|
||||
u = result.get("usage") or {}
|
||||
u.setdefault("provider", "jimeng")
|
||||
u.setdefault("model_key", "jimeng-3.0")
|
||||
result["usage"] = u
|
||||
logger.info("%s 即梦生成成功: %s", tag, result.get("video_path"))
|
||||
return result
|
||||
|
||||
def get_last_video_error(self) -> dict:
|
||||
"""返回最近一次 video_generation 失败的详细错误。空 dict 表示上次成功或未调用。"""
|
||||
return dict(self.last_video_error or {})
|
||||
|
||||
|
||||
def json_safe_dumps(obj: Any, max_len: int = 2000) -> str:
|
||||
"""安全 json 序列化,失败则 fallback 到 repr,超长截断。"""
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
s = _json.dumps(obj, ensure_ascii=False, default=str)
|
||||
except Exception:
|
||||
s = repr(obj)
|
||||
if len(s) > max_len:
|
||||
s = s[:max_len] + f"...(truncated, total {len(s)})"
|
||||
return s
|
||||
|
||||
|
||||
# ── 单例 ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -617,24 +617,19 @@ def call_video_generation(
|
||||
reference_images: list[str] | None = None,
|
||||
reference_audios: list[str] | None = None,
|
||||
reference_videos: list[str] | None = None,
|
||||
) -> dict | None:
|
||||
"""调用 Seedance / Wan 视频生成(v1.6.2 多模型版)。
|
||||
) -> str | None:
|
||||
"""调用 Seedance 2.5 生成视频(v1.6.1 单次出片版),返回本地 MP4 路径;失败返回 None。
|
||||
|
||||
成功返回 {"video_path": str, "usage": dict | None}(usage 含 completion_tokens),失败返回 None。
|
||||
失败时错误详情会写入 client.last_video_error,可通过 get_last_video_error() 读取:
|
||||
{"error_code": str, "user_message": str, "status_code": int, "detail": str, ...}
|
||||
v1.6.1 关键约束(避免 20min 卡死):
|
||||
- 参考音频/视频/多图全部放进 content 数组并带 role=reference_audio/reference_video/reference_image;
|
||||
- 纯首帧无参考时(first_frame 模式),Seedance 2.5 强制 ratio=adaptive;
|
||||
传了参考音/视/多图时走 omni_reference 模式,ratio 可指定为 9:16(客户端内部自动判断)。
|
||||
- ratio 默认 9:16(竖屏),客户端会根据是否有参考自动在 first_frame/adaptive 与 omni/9:16 间切换;
|
||||
若创建任务因 ratio 报错(HTTP 400),客户端会自动回退到 adaptive 再试一次。
|
||||
"""
|
||||
client = get_doubao_client()
|
||||
if not client.is_available:
|
||||
msg = "豆包客户端未配置(DOUBAO_API_KEY 缺失),跳过视频生成"
|
||||
logger.warning("[ai_service] %s", msg)
|
||||
# 写入 last_video_error 供上层读取
|
||||
client.last_video_error = {
|
||||
"error_code": "auth_error",
|
||||
"user_message": "视频生成服务未配置,请联系管理员。",
|
||||
"status_code": 0,
|
||||
"detail": msg,
|
||||
}
|
||||
logger.warning("[ai_service] 豆包客户端未配置,跳过视频生成")
|
||||
return None
|
||||
effective_ratio = ratio or "9:16"
|
||||
try:
|
||||
@@ -656,21 +651,4 @@ def call_video_generation(
|
||||
return client.video_generation(**kwargs)
|
||||
except Exception as e:
|
||||
logger.error("[ai_service] call_video_generation 异常: %s", e, exc_info=True)
|
||||
client.last_video_error = {
|
||||
"error_code": "unknown",
|
||||
"user_message": f"视频生成异常:{e!s}"[:200],
|
||||
"status_code": 0,
|
||||
"detail": str(e),
|
||||
}
|
||||
return None
|
||||
|
||||
|
||||
def get_last_video_error() -> dict:
|
||||
"""读取最近一次视频生成失败的详细错误(含 error_code/user_message/status_code/detail)。
|
||||
成功或未调用过返回空 dict。
|
||||
"""
|
||||
try:
|
||||
client = get_doubao_client()
|
||||
return client.get_last_video_error() if hasattr(client, "get_last_video_error") else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
@@ -1,344 +0,0 @@
|
||||
"""DashScope 客户端(阿里云百炼 Wan 3.0 等非方舟模型)。
|
||||
|
||||
#2159: 新增 Wan 3.0 视频生成支持。DashScope 异步协议:
|
||||
- POST {base_url}/services/aigc/video-generation/video-synthesis (X-DashScope-Async: enable)
|
||||
→ 返回 output.task_id
|
||||
- GET {base_url}/tasks/{task_id} 轮询状态
|
||||
→ SUCCEEDED 时 output.video_url 可下载
|
||||
认证:Authorization: Bearer {DASHSCOPE_API_KEY}
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
# 网络/超时类异常父类集合:覆盖 Timeout/Connect/Network/ReadTimeout/WriteTimeout/PoolTimeout
|
||||
_HTTP_NETWORK_ERRORS = ()
|
||||
try:
|
||||
_HTTP_NETWORK_ERRORS = (httpx.TimeoutException, httpx.NetworkError)
|
||||
except Exception:
|
||||
_HTTP_NETWORK_ERRORS = (Exception,)
|
||||
|
||||
_HTTP_STATUS_ERROR = httpx.HTTPStatusError if hasattr(httpx, "HTTPStatusError") else Exception
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_DASHSCOPE_CLIENT_SINGLETON: "DashScopeClient | None" = None
|
||||
|
||||
|
||||
def _classify_dashscope_error(status_code: int, body: str, task_msg: str = "") -> tuple[str, str]:
|
||||
"""DashScope 错误分类,返回 (error_code, user_message)。"""
|
||||
body_lower = (body or "").lower()
|
||||
msg_in_body = task_msg or ""
|
||||
try:
|
||||
import json as _json
|
||||
|
||||
parsed = _json.loads(body or "{}")
|
||||
if isinstance(parsed, dict):
|
||||
msg_in_body = msg_in_body or str(parsed.get("message", "") or "")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if status_code in (401, 403):
|
||||
return "auth_error", "Wan 3.0 服务鉴权失败(DASHSCOPE_API_KEY 无效或过期),请联系管理员。"
|
||||
if status_code == 429 or "rate" in body_lower or "throttl" in body_lower:
|
||||
return "rate_limit", "Wan 3.0 服务繁忙(限流),请稍等1-2分钟后重试。"
|
||||
if status_code == 400 and any(
|
||||
kw in body_lower for kw in ("portrait", "真人", "人脸", "肖像", "content_violation", "risk", "blocked")
|
||||
):
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"参考素材包含真人照片或违规内容被安全策略拦截,请移除真人图片或调整文案后重试。",
|
||||
)
|
||||
if status_code == 404 or ("not found" in body_lower) or ("model" in body_lower and "not exist" in body_lower):
|
||||
return "model_not_found", "Wan 3.0 模型未开通或模型ID无效,请联系管理员。"
|
||||
if status_code in (402, 400) and ("quota" in body_lower or "billing" in body_lower or "insufficient" in body_lower):
|
||||
return "quota_exceeded", "Wan 3.0 服务配额不足,请联系管理员充值或稍后重试。"
|
||||
if status_code == 400:
|
||||
return "invalid_param", f"Wan 3.0 参数错误:{msg_in_body or body[:200]}"
|
||||
if status_code == 0:
|
||||
return "network_error", "Wan 3.0 服务连接失败(网络超时),请稍后重试。"
|
||||
# 任务内失败
|
||||
if task_msg and any(kw in task_msg.lower() for kw in ("portrait", "真人", "人脸", "violation", "blocked")):
|
||||
return "portrait_intercept", "Wan 3.0 视频内容被安全策略拦截,请调整文案或参考图后重试。"
|
||||
detail = msg_in_body or body[:200]
|
||||
return "unknown", f"Wan 3.0 视频生成失败(HTTP {status_code}):{detail}"
|
||||
|
||||
|
||||
class DashScopeClient:
|
||||
"""阿里云 DashScope 异步 API 客户端(Wan 3.0 等视频生成)。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
settings = get_shared_settings()
|
||||
self.api_key: str = getattr(settings, "dashscope_api_key", "") or os.getenv("DASHSCOPE_API_KEY", "")
|
||||
self.base_url: str = (
|
||||
getattr(settings, "dashscope_base_url", "") or "https://dashscope.aliyuncs.com/api/v1"
|
||||
).rstrip("/")
|
||||
self.poll_interval: int = int(getattr(settings, "dashscope_video_poll_interval", 10) or 10)
|
||||
self.total_timeout: int = int(getattr(settings, "dashscope_video_timeout", 900) or 900)
|
||||
self.max_retries: int = 2
|
||||
self.last_video_error: dict = {}
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.api_key)
|
||||
|
||||
def get_last_video_error(self) -> dict:
|
||||
return dict(self.last_video_error or {})
|
||||
|
||||
def _set_error(self, error_code: str, user_message: str, status_code: int = 0, detail: str = "", **extra) -> None:
|
||||
self.last_video_error = {
|
||||
"error_code": error_code,
|
||||
"user_message": user_message,
|
||||
"status_code": status_code,
|
||||
"detail": detail[:500] if detail else "",
|
||||
**extra,
|
||||
}
|
||||
|
||||
def video_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
image_url: str | None = None,
|
||||
duration: int = 5,
|
||||
ratio: str | None = "9:16",
|
||||
resolution: str = "720p",
|
||||
watermark: bool = False,
|
||||
output_dir: str | None = None,
|
||||
model: str = "wan3.0-video",
|
||||
) -> dict | None:
|
||||
"""调用 DashScope 异步视频合成接口,轮询完成后下载到本地。
|
||||
|
||||
返回 {"video_path": str, "usage": dict | None};失败返回 None,错误详情写入 self.last_video_error。
|
||||
"""
|
||||
self.last_video_error = {}
|
||||
if not self.is_available:
|
||||
self._set_error("auth_error", "Wan 3.0 API key 未配置,请联系管理员。", detail="dashscope api_key empty")
|
||||
logger.error("[dashscope] API key 未配置,无法调用视频生成")
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
self._set_error("invalid_param", "视频生成提示词不能为空。", detail="empty prompt")
|
||||
return None
|
||||
|
||||
# DashScope 分辨率参数:720P / 1080P / 480P(大写 P)
|
||||
res_upper = (resolution or "720p").upper().replace("P", "P")
|
||||
if res_upper == "480P":
|
||||
ds_res = "480P"
|
||||
elif res_upper == "1080P":
|
||||
ds_res = "1080P"
|
||||
else:
|
||||
ds_res = "720P"
|
||||
|
||||
# 构造 input+parameters
|
||||
input_obj: dict[str, Any] = {"prompt": prompt.strip()}
|
||||
if image_url:
|
||||
input_obj["img_url"] = image_url
|
||||
params: dict[str, Any] = {
|
||||
"resolution": ds_res,
|
||||
"duration": str(float(duration)),
|
||||
"watermark": bool(watermark),
|
||||
}
|
||||
# 比例透传:Wan 支持 "9:16" / "16:9" / "1:1" 等
|
||||
if ratio and ratio != "adaptive":
|
||||
params["aspect_ratio"] = ratio
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
"model": model,
|
||||
"input": input_obj,
|
||||
"parameters": params,
|
||||
}
|
||||
headers = {
|
||||
"Authorization": f"Bearer {self.api_key}",
|
||||
"Content-Type": "application/json",
|
||||
"X-DashScope-Async": "enable",
|
||||
}
|
||||
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",
|
||||
model,
|
||||
duration,
|
||||
ratio,
|
||||
ds_res,
|
||||
bool(image_url),
|
||||
)
|
||||
logger.info("[dashscope] 创建任务 payload: model=%s params=%s", model, params)
|
||||
|
||||
# 创建任务
|
||||
task_id: str | None = None
|
||||
last_sc = 0
|
||||
last_body = ""
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
resp = httpx.post(create_url, headers=headers, json=payload, timeout=60)
|
||||
sc = int(getattr(resp, "status_code", 0) or 0)
|
||||
body_text = (getattr(resp, "text", "") or "")[:2000]
|
||||
last_sc = sc
|
||||
last_body = body_text
|
||||
if sc >= 400:
|
||||
logger.error("[dashscope] 创建任务 HTTP %d: %s", sc, body_text)
|
||||
if sc >= 500 and attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
err_code, user_msg = _classify_dashscope_error(sc, body_text)
|
||||
self._set_error(err_code, user_msg, sc, body_text, model=model)
|
||||
return None
|
||||
data = resp.json()
|
||||
tid = (data.get("output") or {}).get("task_id")
|
||||
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))
|
||||
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:
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.5 * (2**attempt))
|
||||
continue
|
||||
logger.error("[dashscope] 创建任务最终失败: %s", _e)
|
||||
self._set_error("unknown", f"Wan 3.0 创建任务异常:{_e!s}"[:200], 0, str(_e))
|
||||
return None
|
||||
if not task_id:
|
||||
if not self.last_video_error:
|
||||
err_code, user_msg = _classify_dashscope_error(last_sc, last_body)
|
||||
self._set_error(err_code, user_msg, last_sc, last_body, model=model)
|
||||
return None
|
||||
|
||||
# 轮询任务
|
||||
poll_url = f"{self.base_url}/tasks/{task_id}"
|
||||
deadline = time.time() + self.total_timeout
|
||||
video_url: str | None = None
|
||||
usage: dict | None = None
|
||||
poll_count = 0
|
||||
last_status = ""
|
||||
while time.time() < deadline:
|
||||
poll_count += 1
|
||||
try:
|
||||
r = httpx.get(poll_url, headers=headers, timeout=30)
|
||||
psc = int(getattr(r, "status_code", 0) or 0)
|
||||
pbody = (getattr(r, "text", "") or "")[:1500]
|
||||
if psc >= 400:
|
||||
logger.warning("[dashscope] 轮询 HTTP %d: %s", psc, pbody[:300])
|
||||
if poll_count < 3:
|
||||
time.sleep(self.poll_interval)
|
||||
continue
|
||||
err_code, user_msg = _classify_dashscope_error(psc, pbody)
|
||||
self._set_error(err_code, user_msg, psc, pbody, task_id=task_id)
|
||||
return None
|
||||
d = r.json()
|
||||
out = d.get("output") or {}
|
||||
task_status = out.get("task_status") or d.get("task_status") or ""
|
||||
last_status = task_status
|
||||
if task_status == "SUCCEEDED":
|
||||
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")
|
||||
if video_url:
|
||||
logger.info("[dashscope] 任务 %s 完成: %s", task_id, video_url[:120])
|
||||
break
|
||||
logger.error("[dashscope] 任务 %s SUCCEEDED 但无 video_url: %s", task_id, str(d)[:500])
|
||||
self._set_error(
|
||||
"unknown",
|
||||
"Wan 3.0 任务成功但未返回视频URL,请联系管理员。",
|
||||
200,
|
||||
str(d)[:500],
|
||||
task_id=task_id,
|
||||
)
|
||||
return None
|
||||
if task_status in ("FAILED", "FAILED_WITH_ERROR", "ERROR"):
|
||||
msg = out.get("message") or d.get("message") or out.get("error_msg") or "unknown error"
|
||||
logger.error("[dashscope] 任务 %s 失败: %s", task_id, msg)
|
||||
err_code, user_msg = _classify_dashscope_error(200, "", msg)
|
||||
self._set_error(err_code, user_msg, 200, msg, task_id=task_id, last_status=task_status)
|
||||
return None
|
||||
if task_status in ("CANCELED", "CANCELLED"):
|
||||
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:
|
||||
logger.warning("[dashscope] 轮询异常: %s", e)
|
||||
time.sleep(self.poll_interval)
|
||||
if not video_url:
|
||||
logger.error("[dashscope] 任务 %s 轮询超时(%ds)", task_id, self.total_timeout)
|
||||
self._set_error(
|
||||
"network_error",
|
||||
f"Wan 3.0 视频生成超时(>{self.total_timeout}s),任务仍在排队,请稍后重试。",
|
||||
0,
|
||||
f"timeout after {self.total_timeout}s, polls={poll_count}, last_status={last_status}",
|
||||
task_id=task_id,
|
||||
last_status=last_status,
|
||||
)
|
||||
return None
|
||||
|
||||
# 下载视频
|
||||
out_dir = output_dir or os.path.join(os.getcwd(), "seedance_outputs")
|
||||
os.makedirs(out_dir, exist_ok=True)
|
||||
suffix = Path(urlparse(video_url).path).suffix or ".mp4"
|
||||
if suffix.lower() not in (".mp4", ".mov", ".webm"):
|
||||
suffix = ".mp4"
|
||||
safe_tid = "".join(c if c.isalnum() or c in "-_" else "_" for c in task_id)[:40]
|
||||
out_path = os.path.join(out_dir, f"wan_{safe_tid}{suffix}")
|
||||
try:
|
||||
with httpx.stream("GET", video_url, timeout=300, follow_redirects=True) as resp:
|
||||
dsc = int(getattr(resp, "status_code", 0) or 0)
|
||||
if dsc >= 400:
|
||||
logger.error("[dashscope] 下载 HTTP %d", dsc)
|
||||
self._set_error("network_error", "Wan 3.0 视频下载失败(HTTP错误),请稍后重试。", dsc)
|
||||
return None
|
||||
with open(out_path, "wb") as f:
|
||||
for chunk in resp.iter_bytes(chunk_size=1024 * 256):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
except Exception as e:
|
||||
logger.error("[dashscope] 下载视频失败: %s", e, exc_info=True)
|
||||
self._set_error("network_error", f"Wan 3.0 视频下载失败:{e!s}"[:200], 0, str(e))
|
||||
return None
|
||||
size = os.path.getsize(out_path) if os.path.exists(out_path) else 0
|
||||
if size < 1024:
|
||||
logger.error("[dashscope] 下载文件过小: %d bytes", size)
|
||||
self._set_error("unknown", "Wan 3.0 视频下载文件过小,请稍后重试。", 0, f"downloaded only {size} bytes")
|
||||
return None
|
||||
logger.info("[dashscope] 视频已下载: %s (%d bytes)", out_path, size)
|
||||
return {"video_path": out_path, "usage": usage}
|
||||
|
||||
|
||||
def get_dashscope_client() -> DashScopeClient | None:
|
||||
"""返回 DashScope 客户端单例;未配置 API key 时返回 None。"""
|
||||
global _DASHSCOPE_CLIENT_SINGLETON
|
||||
if _DASHSCOPE_CLIENT_SINGLETON is None:
|
||||
_DASHSCOPE_CLIENT_SINGLETON = DashScopeClient()
|
||||
if not _DASHSCOPE_CLIENT_SINGLETON.is_available:
|
||||
return None
|
||||
return _DASHSCOPE_CLIENT_SINGLETON
|
||||
@@ -1,526 +0,0 @@
|
||||
"""即梦(Jimeng)视觉 API 客户端 —— 火山引擎 cvtob。
|
||||
|
||||
#2169: 真人参考图在方舟 Seedance 走 B 端审核会被 50411 拦截,
|
||||
即梦走 C 端审核链路,普通真人照片可过审。接入即梦 i2v 作为参考图场景兜底通道。
|
||||
|
||||
接口协议(jimeng_i2v_first_v30 —— 视频3.0 720P 首帧图生视频):
|
||||
- 接口地址:https://visual.volcengineapi.com
|
||||
- 鉴权:火山 V4 签名(Region=cn-north-1, Service=cv),使用 AK/SK
|
||||
- 提交任务:POST ?Action=CVSync2AsyncSubmitTask&Version=2022-08-31
|
||||
body: {"req_key": "jimeng_i2v_first_v30", "image_urls": ["<url>"], "prompt": "...", "seed": -1, "frames": 121}
|
||||
-> {"code": 10000, "data": {"task_id": "..."}}
|
||||
- 查询任务:POST ?Action=CVSync2AsyncGetResult&Version=2022-08-31
|
||||
body: {"req_key": "jimeng_i2v_first_v30", "task_id": "..."}
|
||||
-> {"code": 10000, "data": {"status": "in_queue|generating|done", "video_url": "..."}}
|
||||
- 视频 URL 有效期 1 小时,必须立即下载到本地。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import time
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from urllib.parse import quote, urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
from packages.shared.config import get_shared_settings
|
||||
|
||||
# 网络/超时类异常父类集合
|
||||
_HTTP_NETWORK_ERRORS = ()
|
||||
try:
|
||||
_HTTP_NETWORK_ERRORS = (httpx.TimeoutException, httpx.NetworkError)
|
||||
except Exception:
|
||||
_HTTP_NETWORK_ERRORS = (Exception,)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_JIMENG_CLIENT_SINGLETON: "JimengClient | None" = None
|
||||
|
||||
# ── V4 签名常量 ────────────────────────────────────────────────────────
|
||||
_JIMENG_REGION = "cn-north-1"
|
||||
_JIMENG_SERVICE = "cv"
|
||||
_JIMENG_VERSION = "2022-08-31"
|
||||
_ACTION_SUBMIT = "CVSync2AsyncSubmitTask"
|
||||
_ACTION_POLL = "CVSync2AsyncGetResult"
|
||||
_CONTENT_TYPE = "application/json"
|
||||
_SIGNED_HEADERS_LIST = ["content-type", "host", "x-content-sha256", "x-date"]
|
||||
_SIGNED_HEADERS_STR = ";".join(_SIGNED_HEADERS_LIST)
|
||||
|
||||
|
||||
def _norm_query(params: dict[str, str]) -> str:
|
||||
"""构造规范查询串:按 key 排序,URL 编码(safe=-_.~),空格->%20。"""
|
||||
parts = []
|
||||
for k in sorted(params.keys()):
|
||||
v = params[k]
|
||||
ek = quote(str(k), safe="-_.~")
|
||||
ev = quote(str(v), safe="-_.~").replace("+", "%20")
|
||||
parts.append(f"{ek}={ev}")
|
||||
return "&".join(parts)
|
||||
|
||||
|
||||
def _hmac_sha256(key: bytes, msg: str) -> bytes:
|
||||
return hmac.new(key, msg.encode("utf-8"), hashlib.sha256).digest()
|
||||
|
||||
|
||||
def _sha256_hex(data: bytes) -> str:
|
||||
return hashlib.sha256(data).hexdigest()
|
||||
|
||||
|
||||
def _sign_v4(
|
||||
ak: str,
|
||||
sk: str,
|
||||
method: str,
|
||||
host: str,
|
||||
query: dict[str, str],
|
||||
body_bytes: bytes,
|
||||
x_date: str,
|
||||
) -> dict[str, str]:
|
||||
"""火山 V4 签名,返回需要附加到请求的 headers 字典。
|
||||
|
||||
x_date 形如 "20260101T120000Z"(UTC)。
|
||||
short_date = x_date[:8](YYYYMMDD)。
|
||||
"""
|
||||
short_date = x_date[:8]
|
||||
payload_hash = _sha256_hex(body_bytes)
|
||||
canon_uri = "/"
|
||||
canon_query = _norm_query(query)
|
||||
canon_headers = f"content-type:{_CONTENT_TYPE}\nhost:{host}\nx-content-sha256:{payload_hash}\nx-date:{x_date}\n"
|
||||
canon_request = f"{method}\n{canon_uri}\n{canon_query}\n{canon_headers}\n{_SIGNED_HEADERS_STR}\n{payload_hash}"
|
||||
credential_scope = f"{short_date}/{_JIMENG_REGION}/{_JIMENG_SERVICE}/request"
|
||||
string_to_sign = f"HMAC-SHA256\n{x_date}\n{credential_scope}\n{_sha256_hex(canon_request.encode('utf-8'))}"
|
||||
k_date = _hmac_sha256(sk.encode("utf-8"), short_date)
|
||||
k_region = _hmac_sha256(k_date, _JIMENG_REGION)
|
||||
k_service = _hmac_sha256(k_region, _JIMENG_SERVICE)
|
||||
k_signing = _hmac_sha256(k_service, "request")
|
||||
signature = hmac.new(k_signing, string_to_sign.encode("utf-8"), hashlib.sha256).hexdigest()
|
||||
authorization = (
|
||||
f"HMAC-SHA256 Credential={ak}/{credential_scope}, SignedHeaders={_SIGNED_HEADERS_STR}, Signature={signature}"
|
||||
)
|
||||
return {
|
||||
"Content-Type": _CONTENT_TYPE,
|
||||
"Host": host,
|
||||
"X-Content-Sha256": payload_hash,
|
||||
"X-Date": x_date,
|
||||
"Authorization": authorization,
|
||||
}
|
||||
|
||||
|
||||
# ── 错误分类 ──────────────────────────────────────────────────────────
|
||||
|
||||
# 即梦业务码 -> 是否可重试映射
|
||||
_JIMENG_RETRYABLE_CODES = {50511, 50516, 50429, 50430, 50500, 50501}
|
||||
_JIMENG_NON_RETRYABLE_CODES = {50411, 50412, 50413, 50512, 50513, 50514}
|
||||
|
||||
|
||||
def _classify_jimeng_error(status_code: int, body: str, biz_code: int | None = None) -> tuple[str, str, bool]:
|
||||
"""即梦错误分类,返回 (error_code, user_message, is_retryable)。"""
|
||||
code = biz_code if biz_code is not None else 0
|
||||
body_lower = (body or "").lower()
|
||||
|
||||
# 业务码优先
|
||||
if code == 50411:
|
||||
return (
|
||||
"portrait_intercept",
|
||||
"即梦通道:参考图片前审核未通过(Pre Img Risk Not Pass),请更换参考图后重试。",
|
||||
False,
|
||||
)
|
||||
if code == 50511:
|
||||
return "task_failed", "即梦通道:输出图片后审核未通过,可稍后重试。", True
|
||||
if code in (50412, 50413, 50512):
|
||||
return "invalid_param", "即梦通道:提示词或文本审核不通过,请调整文案后重试。", False
|
||||
if code == 50516:
|
||||
return "task_failed", "即梦通道:输出视频后审核未通过,可稍后重试。", True
|
||||
if code in (50429, 50430):
|
||||
return "rate_limit", "即梦通道:QPS/并发超限,请稍等 1-2 分钟后重试。", True
|
||||
if code in (50500, 50501):
|
||||
return "network_error", "即梦通道:服务内部错误,可稍后重试。", True
|
||||
|
||||
# HTTP 层兜底
|
||||
if status_code in (401, 403):
|
||||
return "auth_error", "即梦通道:AK/SK 鉴权失败,请联系管理员检查 JIMENG_AK/SK 配置。", False
|
||||
if status_code == 429:
|
||||
return "rate_limit", "即梦通道:服务限流,请稍后重试。", True
|
||||
if status_code == 404:
|
||||
return "model_not_found", "即梦通道:接口不存在(req_key 或 Action 错误),请联系管理员。", False
|
||||
if status_code in (402, 400) and any(kw in body_lower for kw in ("quota", "billing", "insufficient", "余额")):
|
||||
return "quota_exceeded", "即梦通道:账户余额/配额不足,请联系管理员充值。", False
|
||||
if status_code == 400:
|
||||
msg = ""
|
||||
try:
|
||||
msg = str(json.loads(body or "{}").get("message", "") or "")
|
||||
except Exception:
|
||||
pass
|
||||
return "invalid_param", f"即梦通道:参数错误:{msg or body[:200]}", False
|
||||
if status_code == 0:
|
||||
return "network_error", "即梦通道:网络连接失败,请稍后重试。", True
|
||||
# 任务内失败
|
||||
if code and code != 10000:
|
||||
return "unknown", f"即梦通道:视频生成失败(错误码 {code}),请稍后重试。", code in _JIMENG_RETRYABLE_CODES
|
||||
detail = body[:200]
|
||||
return "unknown", f"即梦通道:视频生成失败(HTTP {status_code}):{detail}", False
|
||||
|
||||
|
||||
# ── 即梦客户端 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class JimengClient:
|
||||
"""火山引擎即梦视觉 API(cvtob)异步客户端,支持图生视频首帧(jimeng_i2v_first_v30)。"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
settings = get_shared_settings()
|
||||
self.ak: str = getattr(settings, "jimeng_ak", "") or os.getenv("JIMENG_AK", "")
|
||||
self.sk: str = getattr(settings, "jimeng_sk", "") or os.getenv("JIMENG_SK", "")
|
||||
self.base_url: str = (getattr(settings, "jimeng_base_url", "") or "https://visual.volcengineapi.com").rstrip(
|
||||
"/"
|
||||
)
|
||||
self.req_key: str = getattr(settings, "jimeng_req_key", "") or "jimeng_i2v_first_v30"
|
||||
self.poll_interval: int = int(getattr(settings, "jimeng_video_poll_interval", 5) or 5)
|
||||
self.total_timeout: int = int(getattr(settings, "jimeng_video_timeout", 600) or 600)
|
||||
self.max_retries: int = 2
|
||||
self.last_video_error: dict = {}
|
||||
# 解析 base_url 里的 host(用于签名 Host 头)
|
||||
parsed = urlparse(self.base_url)
|
||||
self.host: str = parsed.netloc or "visual.volcengineapi.com"
|
||||
|
||||
@property
|
||||
def is_available(self) -> bool:
|
||||
return bool(self.ak and self.sk)
|
||||
|
||||
def get_last_video_error(self) -> dict:
|
||||
return dict(self.last_video_error or {})
|
||||
|
||||
def _set_error(self, error_code: str, user_message: str, status_code: int = 0, detail: str = "", **extra) -> None:
|
||||
self.last_video_error = {
|
||||
"error_code": error_code,
|
||||
"user_message": user_message,
|
||||
"status_code": status_code,
|
||||
"detail": detail[:500] if detail else "",
|
||||
"provider": "jimeng",
|
||||
**extra,
|
||||
}
|
||||
|
||||
# ── 内部 HTTP:签名 + 请求 ──────────────────────────────────────
|
||||
|
||||
def _signed_request(
|
||||
self,
|
||||
method: str,
|
||||
action: str,
|
||||
body_obj: dict[str, Any],
|
||||
timeout: float = 60.0,
|
||||
) -> tuple[int, str, dict]:
|
||||
"""发送一次带 V4 签名的请求,返回 (status_code, body_text, parsed_json)。"""
|
||||
body_bytes = json.dumps(body_obj, ensure_ascii=False).encode("utf-8")
|
||||
query = {"Action": action, "Version": _JIMENG_VERSION}
|
||||
x_date = datetime.now(timezone.utc).strftime("%Y%m%dT%H%M%SZ")
|
||||
headers = _sign_v4(self.ak, self.sk, method, self.host, query, body_bytes, x_date)
|
||||
url = f"{self.base_url}/?{_norm_query(query)}"
|
||||
resp = httpx.request(
|
||||
method,
|
||||
url,
|
||||
headers=headers,
|
||||
content=body_bytes,
|
||||
timeout=timeout,
|
||||
)
|
||||
sc = int(getattr(resp, "status_code", 0) or 0)
|
||||
text = getattr(resp, "text", "") or ""
|
||||
try:
|
||||
data = resp.json()
|
||||
except Exception:
|
||||
data = {}
|
||||
return sc, text, data
|
||||
|
||||
# ── 提交任务 ────────────────────────────────────────────────────
|
||||
|
||||
def _submit_task(
|
||||
self,
|
||||
prompt: str,
|
||||
image_url: str,
|
||||
frames: int = 121,
|
||||
seed: int = -1,
|
||||
) -> str | None:
|
||||
"""提交图生视频任务,成功返回 task_id;失败写 last_video_error 并返回 None。"""
|
||||
body: dict[str, Any] = {
|
||||
"req_key": self.req_key,
|
||||
"prompt": prompt.strip()[:800],
|
||||
"image_urls": [image_url],
|
||||
"seed": int(seed) if seed and seed > 0 else -1,
|
||||
"frames": int(frames),
|
||||
}
|
||||
last_sc = 0
|
||||
last_body = ""
|
||||
for attempt in range(self.max_retries + 1):
|
||||
try:
|
||||
sc, text, data = self._signed_request("POST", _ACTION_SUBMIT, body, timeout=60.0)
|
||||
last_sc, last_body = sc, text
|
||||
if sc >= 400:
|
||||
logger.error("[jimeng] 提交 HTTP %d: %s", sc, text[:500])
|
||||
if sc >= 500 and attempt < self.max_retries:
|
||||
time.sleep(0.8 * (2**attempt))
|
||||
continue
|
||||
biz_code = data.get("code") if isinstance(data, dict) else None
|
||||
err_code, user_msg, _ = _classify_jimeng_error(sc, text, biz_code)
|
||||
self._set_error(err_code, user_msg, sc, text, req_key=self.req_key)
|
||||
return None
|
||||
code = data.get("code") if isinstance(data, dict) else None
|
||||
if code == 10000:
|
||||
d = data.get("data") or {}
|
||||
tid = d.get("task_id")
|
||||
if tid:
|
||||
return str(tid)
|
||||
err_code, user_msg, retry = _classify_jimeng_error(sc, text, code)
|
||||
logger.error(
|
||||
"[jimeng] 提交业务错误 code=%s msg=%s",
|
||||
code,
|
||||
(data.get("message") if isinstance(data, dict) else ""),
|
||||
)
|
||||
if retry and attempt < self.max_retries:
|
||||
time.sleep(0.8 * (2**attempt))
|
||||
continue
|
||||
self._set_error(err_code, user_msg, sc, text, req_key=self.req_key, biz_code=code)
|
||||
return None
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
last_sc, last_body = 0, f"network error: {ne}"
|
||||
logger.warning(
|
||||
"[jimeng] 提交网络异常 %s,重试 %d/%d", type(ne).__name__, attempt + 1, self.max_retries + 1
|
||||
)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.8 * (2**attempt))
|
||||
continue
|
||||
self._set_error("network_error", "即梦通道:提交任务网络异常,请稍后重试。", 0, str(ne))
|
||||
return None
|
||||
except Exception as e:
|
||||
last_sc, last_body = 0, f"exception: {e}"
|
||||
logger.error("[jimeng] 提交异常: %s", e, exc_info=True)
|
||||
if attempt < self.max_retries:
|
||||
time.sleep(0.8 * (2**attempt))
|
||||
continue
|
||||
self._set_error("unknown", f"即梦通道:提交任务异常:{e!s}"[:200], 0, str(e))
|
||||
return None
|
||||
if not self.last_video_error:
|
||||
err_code, user_msg, _ = _classify_jimeng_error(last_sc, last_body)
|
||||
self._set_error(err_code, user_msg, last_sc, last_body)
|
||||
return None
|
||||
|
||||
# ── 轮询结果 ────────────────────────────────────────────────────
|
||||
|
||||
def _poll_result(self, task_id: str) -> str | None:
|
||||
"""轮询任务直到 done/failed/expired/timeout,成功返回 video_url。"""
|
||||
deadline = time.time() + self.total_timeout
|
||||
poll_count = 0
|
||||
last_status = ""
|
||||
poll_body = {"req_key": self.req_key, "task_id": task_id}
|
||||
while time.time() < deadline:
|
||||
poll_count += 1
|
||||
try:
|
||||
sc, text, data = self._signed_request("POST", _ACTION_POLL, poll_body, timeout=30.0)
|
||||
if sc >= 400:
|
||||
logger.warning("[jimeng] 轮询 HTTP %d: %s", sc, text[:300])
|
||||
if poll_count < 3:
|
||||
time.sleep(self.poll_interval)
|
||||
continue
|
||||
err_code, user_msg, _ = _classify_jimeng_error(sc, text)
|
||||
self._set_error(err_code, user_msg, sc, text, task_id=task_id)
|
||||
return None
|
||||
code = data.get("code") if isinstance(data, dict) else None
|
||||
d = data.get("data") if isinstance(data, dict) else None
|
||||
if code != 10000 or not isinstance(d, dict):
|
||||
err_code, user_msg, retry = _classify_jimeng_error(sc, text, code)
|
||||
logger.error(
|
||||
"[jimeng] 轮询业务错误 task=%s code=%s msg=%s",
|
||||
task_id,
|
||||
code,
|
||||
(data.get("message") if isinstance(data, dict) else ""),
|
||||
)
|
||||
if retry and poll_count < 3:
|
||||
time.sleep(self.poll_interval)
|
||||
continue
|
||||
self._set_error(err_code, user_msg, sc, text, task_id=task_id, biz_code=code)
|
||||
return None
|
||||
status = d.get("status", "") or ""
|
||||
last_status = status
|
||||
if status == "done":
|
||||
video_url = d.get("video_url") or ""
|
||||
if video_url:
|
||||
logger.info("[jimeng] 任务 %s 完成 polls=%d", task_id, poll_count)
|
||||
return str(video_url)
|
||||
logger.error("[jimeng] 任务 %s done 但无 video_url: %s", task_id, str(d)[:500])
|
||||
self._set_error(
|
||||
"unknown",
|
||||
"即梦通道:任务成功但未返回视频URL,请联系管理员。",
|
||||
200,
|
||||
str(d)[:500],
|
||||
task_id=task_id,
|
||||
)
|
||||
return None
|
||||
if status in ("not_found", "expired"):
|
||||
logger.error("[jimeng] 任务 %s 状态 %s", task_id, status)
|
||||
self._set_error(
|
||||
"network_error" if status == "expired" else "unknown",
|
||||
f"即梦通道:任务{'已过期' if status == 'expired' else '未找到'},请重新提交。",
|
||||
200,
|
||||
f"task {status}",
|
||||
task_id=task_id,
|
||||
)
|
||||
return None
|
||||
if poll_count % 6 == 0:
|
||||
logger.info("[jimeng] 轮询中 task=%s status=%s polls=%d", task_id, status, poll_count)
|
||||
except _HTTP_NETWORK_ERRORS as ne:
|
||||
logger.warning("[jimeng] 轮询网络异常 %s", ne)
|
||||
except Exception as e:
|
||||
logger.debug("[jimeng] 轮询异常: %s", e)
|
||||
time.sleep(self.poll_interval)
|
||||
logger.error(
|
||||
"[jimeng] 任务 %s 轮询超时(%ds)polls=%d last_status=%s",
|
||||
task_id,
|
||||
self.total_timeout,
|
||||
poll_count,
|
||||
last_status,
|
||||
)
|
||||
self._set_error(
|
||||
"network_error",
|
||||
f"即梦通道:视频生成超时(>{self.total_timeout}s),任务仍在排队,请稍后重试。",
|
||||
0,
|
||||
f"timeout after {self.total_timeout}s, polls={poll_count}, last_status={last_status}",
|
||||
task_id=task_id,
|
||||
last_status=last_status,
|
||||
)
|
||||
return None
|
||||
|
||||
# ── 下载视频 ────────────────────────────────────────────────────
|
||||
|
||||
def _download_video(self, video_url: str, output_dir: str, task_id: str) -> str | None:
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
suffix = Path(urlparse(video_url).path).suffix or ".mp4"
|
||||
if suffix.lower() not in (".mp4", ".mov", ".webm"):
|
||||
suffix = ".mp4"
|
||||
safe_tid = "".join(c if c.isalnum() or c in "-_" else "_" for c in task_id)[:40]
|
||||
out_path = os.path.join(output_dir, f"jimeng_{safe_tid}_{uuid.uuid4().hex[:8]}{suffix}")
|
||||
try:
|
||||
with httpx.stream("GET", video_url, timeout=180, follow_redirects=True) as resp:
|
||||
dsc = int(getattr(resp, "status_code", 0) or 0)
|
||||
if dsc >= 400:
|
||||
logger.error("[jimeng] 下载 HTTP %d", dsc)
|
||||
self._set_error("network_error", "即梦通道:视频下载失败(HTTP错误),请稍后重试。", dsc)
|
||||
return None
|
||||
with open(out_path, "wb") as f:
|
||||
for chunk in resp.iter_bytes(chunk_size=1024 * 256):
|
||||
if chunk:
|
||||
f.write(chunk)
|
||||
except Exception as e:
|
||||
logger.error("[jimeng] 下载视频失败: %s", e, exc_info=True)
|
||||
self._set_error("network_error", f"即梦通道:视频下载失败:{e!s}"[:200], 0, str(e))
|
||||
return None
|
||||
size = os.path.getsize(out_path) if os.path.exists(out_path) else 0
|
||||
if size < 1024:
|
||||
logger.error("[jimeng] 下载文件过小: %d bytes", size)
|
||||
self._set_error("unknown", "即梦通道:视频下载文件过小,请稍后重试。", 0, f"downloaded only {size} bytes")
|
||||
try:
|
||||
os.remove(out_path)
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
logger.info("[jimeng] 视频已下载: %s (%d bytes)", out_path, size)
|
||||
return out_path
|
||||
|
||||
# ── 对外主入口 ──────────────────────────────────────────────────
|
||||
|
||||
def video_generation(
|
||||
self,
|
||||
prompt: str,
|
||||
*,
|
||||
image_url: str,
|
||||
duration: int = 5,
|
||||
ratio: str | None = "9:16",
|
||||
resolution: str = "720p",
|
||||
output_dir: str | None = None,
|
||||
generate_audio: bool = False,
|
||||
) -> dict | None:
|
||||
"""即梦图生视频主入口。
|
||||
|
||||
成功返回 {"video_path": str, "usage": {"provider","duration_seconds","frames","req_key","billing_mode"}};
|
||||
失败返回 None,详情在 self.last_video_error。
|
||||
|
||||
注意:jimeng_i2v_first_v30 不支持原生音频(generate_audio 被忽略,返回无声视频),
|
||||
音频由后续 ffmpeg 合成阶段叠加 TTS。
|
||||
支持时长:5s(frames=121)/10s(frames=241),>10s 截断并打 warning。
|
||||
分辨率固定 720P;ratio 对首帧 i2v 无效(自动按图片比例)。
|
||||
"""
|
||||
self.last_video_error = {}
|
||||
if not self.is_available:
|
||||
self._set_error(
|
||||
"auth_error",
|
||||
"即梦通道未配置(JIMENG_AK/SK 缺失),请联系管理员。",
|
||||
detail="jimeng ak/sk empty",
|
||||
)
|
||||
logger.error("[jimeng] AK/SK 未配置,无法调用")
|
||||
return None
|
||||
if not prompt or not prompt.strip():
|
||||
self._set_error("invalid_param", "视频生成提示词不能为空。", detail="empty prompt")
|
||||
return None
|
||||
if not image_url or not image_url.strip():
|
||||
self._set_error("invalid_param", "即梦图生视频必须提供参考图片。", detail="empty image_url")
|
||||
return None
|
||||
|
||||
dur = int(duration or 5)
|
||||
if dur <= 5:
|
||||
frames = 121
|
||||
real_dur = 5
|
||||
elif dur <= 10:
|
||||
frames = 241
|
||||
real_dur = 10
|
||||
else:
|
||||
logger.warning("[jimeng] 请求时长 %ds 超出即梦 i2v 上限 10s,截断到 10s(frames=241)", dur)
|
||||
frames = 241
|
||||
real_dur = 10
|
||||
|
||||
out_dir = output_dir or "/tmp"
|
||||
logger.info(
|
||||
"[jimeng] 提交任务: req_key=%s dur=%ds(frames=%d) ratio=%s res=%s gen_audio=%s img=%s",
|
||||
self.req_key,
|
||||
real_dur,
|
||||
frames,
|
||||
ratio,
|
||||
resolution,
|
||||
generate_audio,
|
||||
bool(image_url),
|
||||
)
|
||||
|
||||
task_id = self._submit_task(prompt=prompt, image_url=image_url, frames=frames, seed=-1)
|
||||
if not task_id:
|
||||
return None
|
||||
logger.info("[jimeng] 任务已提交: task_id=%s", task_id)
|
||||
|
||||
video_url = self._poll_result(task_id)
|
||||
if not video_url:
|
||||
return None
|
||||
|
||||
local_path = self._download_video(video_url, out_dir, task_id)
|
||||
if not local_path:
|
||||
return None
|
||||
|
||||
usage = {
|
||||
"provider": "jimeng",
|
||||
"duration_seconds": real_dur,
|
||||
"frames": frames,
|
||||
"req_key": self.req_key,
|
||||
"billing_mode": "per_second",
|
||||
}
|
||||
return {"video_path": local_path, "usage": usage}
|
||||
|
||||
|
||||
def get_jimeng_client() -> "JimengClient | None":
|
||||
"""返回即梦客户端单例;未配置 AK/SK 时返回 None。"""
|
||||
global _JIMENG_CLIENT_SINGLETON
|
||||
if _JIMENG_CLIENT_SINGLETON is None:
|
||||
_JIMENG_CLIENT_SINGLETON = JimengClient()
|
||||
if not _JIMENG_CLIENT_SINGLETON.is_available:
|
||||
return None
|
||||
return _JIMENG_CLIENT_SINGLETON
|
||||
@@ -377,12 +377,32 @@ class TestPrepareNarrativeVoice:
|
||||
assert ei.value.status_code == 502
|
||||
assert "配音合成失败" in ei.value.message
|
||||
|
||||
def test_no_points_service_invoked(self, monkeypatch):
|
||||
"""v1.6.2: 叙事配音已免费,不再实例化 PointsService / 扣点/退费。"""
|
||||
# 确认 narrative_service 已不再暴露 PointsService
|
||||
assert not hasattr(ns, "PointsService"), "narrative_service 不应再导入 PointsService"
|
||||
def test_points_insufficient_402(self, monkeypatch):
|
||||
class FakePoints:
|
||||
def deduct_points(self, *a, **k):
|
||||
return {"success": False, "balance": 0}
|
||||
|
||||
class FakeWorkflow:
|
||||
monkeypatch.setattr(ns, "PointsService", lambda: FakePoints())
|
||||
deps = self._deps(points_enabled=True)
|
||||
with pytest.raises(NarrativeError) as ei:
|
||||
prepare_narrative_voice(**deps)
|
||||
assert ei.value.status_code == 402
|
||||
|
||||
def test_points_refund_on_failure(self, monkeypatch):
|
||||
class FakePoints:
|
||||
def __init__(self):
|
||||
self.refunded = 0
|
||||
|
||||
def deduct_points(self, *a, **k):
|
||||
return {"success": True, "balance": 100}
|
||||
|
||||
def refund_points(self, user_id, amount, source, db, ref_id="", **k):
|
||||
self.refunded += amount
|
||||
|
||||
points = FakePoints()
|
||||
monkeypatch.setattr(ns, "PointsService", lambda: points)
|
||||
|
||||
class FailingWorkflow:
|
||||
def __init__(self, *, repository, cosyvoice_service):
|
||||
pass
|
||||
|
||||
@@ -392,22 +412,11 @@ class TestPrepareNarrativeVoice:
|
||||
def process_synthesis_failure(self, job_id, error):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(ns, "TTSWorkflowService", FakeWorkflow)
|
||||
monkeypatch.setattr(ns, "TTSWorkflowService", FailingWorkflow)
|
||||
deps = self._deps(points_enabled=True)
|
||||
with pytest.raises(NarrativeError) as ei:
|
||||
with pytest.raises(NarrativeError):
|
||||
prepare_narrative_voice(**deps)
|
||||
# 走 502 业务错误路径,不再退费
|
||||
assert ei.value.status_code == 502
|
||||
|
||||
def test_module_has_no_points_imports(self):
|
||||
"""模块源码不再包含扣点相关符号。"""
|
||||
import inspect
|
||||
|
||||
src = inspect.getsource(ns)
|
||||
assert "PointsService" not in src
|
||||
assert "calculate_points_cost" not in src
|
||||
assert "_points_scene" not in src
|
||||
assert "_POINTS_SCENE" not in src
|
||||
assert points.refunded > 0
|
||||
|
||||
def test_clone_source_resolves_profile(self, monkeypatch):
|
||||
captured = {}
|
||||
|
||||
@@ -87,9 +87,7 @@ _GEN_TASKS_PATH = Path(__file__).resolve().parents[2] / "apps/api/app/api/routes
|
||||
def _load_infer_func():
|
||||
src = _GEN_TASKS_PATH.read_text()
|
||||
start = src.index("# #2035:文案关键词")
|
||||
# 用紧跟 _infer_expected_categories 后的 logger 行作为结束锚点
|
||||
end_marker = "\nlogger = logging.getLogger"
|
||||
end = src.index(end_marker, start)
|
||||
end = src.index("from packages.middleware")
|
||||
code = src[start:end]
|
||||
ns: dict = {}
|
||||
exec(code, ns)
|
||||
|
||||
@@ -1,25 +1,48 @@
|
||||
"""AI 数字人渲染 — v1.6.2 起免费,不扣积分"""
|
||||
"""AI数字人渲染 积分扣点单元测试 (#1895 P2 step 2.6)"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
class TestAiAvatarRenderFree:
|
||||
def test_ai_digital_human_returns_zero_cost(self):
|
||||
import pytest
|
||||
|
||||
import packages.middleware.points_gate as _pg_module
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _enable(monkeypatch):
|
||||
monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True)
|
||||
yield
|
||||
|
||||
|
||||
class TestAiAvatarRenderPoints:
|
||||
def test_ai_digital_human_per_unit(self):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
assert calculate_points_cost("ai_digital_human", is_member=False, duration_minutes=1) == 0
|
||||
assert calculate_points_cost("ai_digital_human", is_member=True, duration_minutes=5) == 0
|
||||
cost = calculate_points_cost("ai_digital_human", is_member=False, duration_minutes=1)
|
||||
assert cost >= 15
|
||||
|
||||
def test_no_points_gate_decorator(self):
|
||||
def test_decorator_attached(self):
|
||||
from app.api.routes.ai_avatar_render import create_render_job
|
||||
|
||||
assert not hasattr(create_render_job, "__wrapped__")
|
||||
assert hasattr(create_render_job, "__wrapped__"), "missing @points_gate"
|
||||
|
||||
def test_module_has_no_points_imports(self):
|
||||
import inspect
|
||||
def test_insufficient_raises_402(self):
|
||||
from app.api.routes.ai_avatar_render import create_render_job
|
||||
from app.schemas.ai_avatar_render import CreateAiAvatarRenderRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from app.api.routes import ai_avatar_render as mod
|
||||
|
||||
src = inspect.getsource(mod)
|
||||
assert "PointsService" not in src
|
||||
assert "points_gate" not in src
|
||||
db = MagicMock()
|
||||
cu = MagicMock()
|
||||
cu.user.id = "u1"
|
||||
cu.user.is_member = False
|
||||
cu.user.member_type = None
|
||||
svc = MagicMock()
|
||||
body = CreateAiAvatarRenderRequest(lipsync_job_id="lip1")
|
||||
with patch("packages.domain.points_service.PointsService") as MS:
|
||||
msvc = MagicMock()
|
||||
msvc.deduct_points.return_value = {"success": False, "balance": 0}
|
||||
MS.return_value = msvc
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
create_render_job(body=body, current_user=cu, svc=svc, db=db)
|
||||
assert ei.value.status_code == 402
|
||||
|
||||
@@ -112,10 +112,10 @@ class TestVideoGenerationHappyPath:
|
||||
resolution="720p",
|
||||
output_dir=str(tmp_path),
|
||||
)
|
||||
assert out is not None and isinstance(out, dict)
|
||||
assert Path(out["video_path"]).exists()
|
||||
assert Path(out["video_path"]).name == "seedance_task-001_abcd1234.mp4"
|
||||
assert Path(out["video_path"]).read_bytes() == b"FAKEMP4DATA"
|
||||
assert out is not None
|
||||
assert Path(out).exists()
|
||||
assert Path(out).name == "seedance_task-001_abcd1234.mp4"
|
||||
assert Path(out).read_bytes() == b"FAKEMP4DATA"
|
||||
assert calls["post"] == 1
|
||||
assert calls["get"] == 1
|
||||
|
||||
@@ -325,8 +325,8 @@ class TestVideoGenerationPollLoop:
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=200, doubao_video_model="seedance"
|
||||
)
|
||||
out = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert out is not None and isinstance(out, dict)
|
||||
assert Path(out["video_path"]).read_bytes() == b"DATA"
|
||||
assert out is not None
|
||||
assert Path(out).read_bytes() == b"DATA"
|
||||
# queued 和 running 各 sleep 一次
|
||||
assert len(sleeps) >= 2
|
||||
|
||||
@@ -373,8 +373,8 @@ class TestVideoGenerationPollLoop:
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=100, doubao_video_model="seedance"
|
||||
)
|
||||
out = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert out is not None and isinstance(out, dict)
|
||||
assert Path(out["video_path"]).exists()
|
||||
assert out is not None
|
||||
assert Path(out).exists()
|
||||
assert poll_calls["n"] == 2
|
||||
|
||||
def test_default_output_dir_and_audio_watermark(self, tmp_path, monkeypatch):
|
||||
@@ -442,8 +442,8 @@ class TestVideoGenerationPollLoop:
|
||||
out = client.video_generation(
|
||||
"p", duration=3, ratio="1:1", resolution="480p", generate_audio=True, watermark=True
|
||||
)
|
||||
assert out is not None and isinstance(out, dict)
|
||||
assert out["video_path"] == "/tmp/seedance_t-default_00000001.mp4"
|
||||
assert out is not None
|
||||
assert "/tmp/seedance_t-default_00000001.mp4" in out
|
||||
assert captured["json"]["generate_audio"] is True
|
||||
assert captured["json"]["watermark"] is True
|
||||
assert captured["json"]["ratio"] == "1:1"
|
||||
@@ -483,180 +483,3 @@ class TestVideoGenerationCancelled:
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
assert client.video_generation("p", output_dir=str(tmp_path)) is None
|
||||
|
||||
|
||||
# ============ #2157 _resolve_video_model_id 模型ID映射单测 ============
|
||||
|
||||
|
||||
class TestResolveVideoModelId:
|
||||
"""覆盖 _resolve_video_model_id 各分支(#2157 P0 修复)。"""
|
||||
|
||||
def _import_target(self):
|
||||
from packages.shared.ai_client import _resolve_video_model_id
|
||||
|
||||
return _resolve_video_model_id
|
||||
|
||||
def test_none_uses_default(self):
|
||||
fn = self._import_target()
|
||||
with patch("packages.shared.ai_client.get_shared_settings") as ms:
|
||||
ms.return_value = MagicMock(doubao_video_model="doubao-seedance-2-5-260628")
|
||||
assert fn(None) == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_empty_uses_default(self):
|
||||
fn = self._import_target()
|
||||
with patch("packages.shared.ai_client.get_shared_settings") as ms:
|
||||
ms.return_value = MagicMock(doubao_video_model="doubao-seedance-2-5-260628")
|
||||
assert fn(" ") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_doubao_prefix_passthrough(self):
|
||||
fn = self._import_target()
|
||||
assert fn("doubao-seedance-2-5-260628") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_ep_prefix_passthrough(self):
|
||||
fn = self._import_target()
|
||||
assert fn("ep-20260721114705-b568m") == "ep-20260721114705-b568m"
|
||||
|
||||
def test_seedance_2_5_alias(self):
|
||||
fn = self._import_target()
|
||||
assert fn("seedance-2.5") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_seedance_2_0_alias(self):
|
||||
fn = self._import_target()
|
||||
assert fn("seedance-2.0") == "doubao-seedance-2-0-260128"
|
||||
|
||||
def test_seedance_2_0_fast_alias(self):
|
||||
fn = self._import_target()
|
||||
assert fn("seedance-2.0-fast") == "doubao-seedance-2-0-fast-260128"
|
||||
|
||||
def test_seedance_2_0_mini_alias(self):
|
||||
fn = self._import_target()
|
||||
assert fn("seedance-2.0-mini") == "doubao-seedance-2-0-mini-260615"
|
||||
|
||||
def test_wan_3_0_returns_dashscope_provider(self):
|
||||
from packages.shared.ai_client import _resolve_video_provider_and_id
|
||||
|
||||
prov, mid, cfg = _resolve_video_provider_and_id("wan-3.0")
|
||||
assert prov == "dashscope"
|
||||
assert mid == "wan3.0-video"
|
||||
assert cfg.get("billing_mode") == "per_second"
|
||||
|
||||
def test_seedance_2_5_uppercase(self):
|
||||
fn = self._import_target()
|
||||
assert fn("Seedance-2.5") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_seedance_dot_normalize(self):
|
||||
fn = self._import_target()
|
||||
# dot 形式 "seedance-2.5" 直接命中 domain config 的 key(与 2-5 同等)
|
||||
assert fn("seedance-2.5") == "doubao-seedance-2-5-260628"
|
||||
|
||||
def test_unknown_model_falls_back_to_default_seedance_2_5(self, caplog):
|
||||
fn = self._import_target()
|
||||
import logging
|
||||
|
||||
# 未知 model key 会通过 get_viral_video_model_config 回落到 seedance-2.5
|
||||
with caplog.at_level(logging.WARNING, logger="shared.ai_client"):
|
||||
assert fn("some-random-model") == "doubao-seedance-2-5-260628"
|
||||
|
||||
|
||||
# ── #2165 详细错误信息和 last_video_error ─────────────────────────
|
||||
|
||||
|
||||
class TestVideoGenerationLastError:
|
||||
def test_create_400_portrait_returns_user_message(self, tmp_path):
|
||||
"""#2169: HTTP 400 + 真人拦截关键词 → 自动尝试即梦兜底;即梦未配时返回 portrait_intercept。"""
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 400
|
||||
create_resp.text = '{"error":{"code":"ContentRisk","message":"Real person face detected in reference image, portrait blocked"}}'
|
||||
create_resp.json.return_value = {"error": {"code": "ContentRisk", "message": "..."}}
|
||||
create_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
|
||||
"bad", request=MagicMock(), response=create_resp
|
||||
)
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
# jimeng 未配置,fallback 后仍返回 portrait_intercept(提示用户需要配置即梦)
|
||||
patch("packages.shared.jimeng_client.get_jimeng_client", return_value=None),
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=1, doubao_video_model="seedance"
|
||||
)
|
||||
result = client.video_generation("p", output_dir=str(tmp_path), image_url="https://img/x.jpg")
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
assert err["error_code"] == "portrait_intercept"
|
||||
# 即梦兜底未启用时提示包含"真人照片"/"即梦"等关键字
|
||||
assert "真人" in err["user_message"] or "肖像" in err["user_message"] or "即梦" in err["user_message"]
|
||||
# 方舟本身 status_code=400(最后一个错误可能被即梦兜底覆盖,但 error_code 不变)
|
||||
assert err["status_code"] in (0, 400)
|
||||
|
||||
def test_create_401_returns_auth_error(self, tmp_path):
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 401
|
||||
create_resp.text = '{"error":{"message":"Unauthorized"}}'
|
||||
create_resp.json.return_value = {"error": {"message": "Unauthorized"}}
|
||||
create_resp.raise_for_status.side_effect = httpx.HTTPStatusError(
|
||||
"auth", request=MagicMock(), response=create_resp
|
||||
)
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=1, doubao_video_model="seedance"
|
||||
)
|
||||
result = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
assert err["error_code"] == "auth_error"
|
||||
assert err["status_code"] == 401
|
||||
|
||||
def test_poll_failed_returns_task_failed_error(self, tmp_path):
|
||||
"""轮询 status=failed 时应记录 task_failed 错误并含 detail。"""
|
||||
client = _make_client(max_retries=0)
|
||||
create_resp = MagicMock()
|
||||
create_resp.status_code = 200
|
||||
create_resp.json.return_value = {"id": "t-fail"}
|
||||
create_resp.raise_for_status = MagicMock()
|
||||
poll_resp = MagicMock()
|
||||
poll_resp.status_code = 200
|
||||
poll_resp.json.return_value = {
|
||||
"status": "failed",
|
||||
"error": {"code": "InvalidParam", "message": "resolution invalid"},
|
||||
}
|
||||
poll_resp.raise_for_status = MagicMock()
|
||||
with (
|
||||
patch("packages.shared.ai_client.httpx.post", return_value=create_resp),
|
||||
patch("packages.shared.ai_client.httpx.get", return_value=poll_resp),
|
||||
patch("packages.shared.ai_client.time.sleep", return_value=None),
|
||||
patch("packages.shared.ai_client.time.time", side_effect=_fake_time_factory()),
|
||||
patch("packages.shared.ai_client.get_shared_settings") as mock_s,
|
||||
):
|
||||
mock_s.return_value = MagicMock(
|
||||
doubao_video_poll_interval=0, doubao_video_timeout=10, doubao_video_model="seedance"
|
||||
)
|
||||
result = client.video_generation("p", output_dir=str(tmp_path))
|
||||
assert result is None
|
||||
err = client.get_last_video_error()
|
||||
assert err["error_code"] == "task_failed"
|
||||
assert "InvalidParam" in err.get("detail", "") or err["status_code"] == 200
|
||||
|
||||
|
||||
class TestAiServiceLastVideoError:
|
||||
def test_call_video_generation_returns_none_sets_error(self):
|
||||
"""失败后 get_last_video_error 应返回结构化错误信息。"""
|
||||
from packages.shared import ai_service
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.last_video_error = {"error_code": "unknown", "user_message": "test"}
|
||||
mock_client.get_last_video_error.return_value = {"error_code": "unknown", "user_message": "test"}
|
||||
mock_client.video_generation.return_value = None
|
||||
with patch("packages.shared.ai_service.get_doubao_client", return_value=mock_client):
|
||||
assert ai_service.call_video_generation("p") is None
|
||||
err = ai_service.get_last_video_error()
|
||||
assert err["error_code"] == "unknown"
|
||||
assert "user_message" in err
|
||||
|
||||
@@ -1,198 +0,0 @@
|
||||
"""catalog 应用服务单测:会员套餐 / 积分包从共享库读取与字段映射。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _clear_cache():
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
admin_catalog._cache.clear()
|
||||
yield
|
||||
admin_catalog._cache.clear()
|
||||
|
||||
|
||||
def _row(**kw):
|
||||
row = MagicMock()
|
||||
for k, v in kw.items():
|
||||
setattr(row, k, v)
|
||||
return row
|
||||
|
||||
|
||||
class TestMembershipPlans:
|
||||
def test_yearly_plan_mapping(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium_yearly",
|
||||
name="高级会员年卡",
|
||||
description="年度订阅",
|
||||
monthly_price=0,
|
||||
yearly_price=399,
|
||||
quotas={"4k": True, "batch_render": True, "credits_per_month": 500},
|
||||
display_order=1,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
assert len(plans) == 1
|
||||
p = plans[0]
|
||||
assert p["plan_id"] == "premium_yearly"
|
||||
assert p["billing_cycle"] == "yearly"
|
||||
assert p["price_cents"] == 39900
|
||||
assert p["monthly_price_cents"] == 3325
|
||||
assert p["duration_days"] == 365
|
||||
assert p["features"]["4K 超清分辨率"] is True
|
||||
assert p["features"]["credits_per_month"] == 500
|
||||
session.close.assert_called_once()
|
||||
|
||||
def test_monthly_plan_mapping(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium_monthly",
|
||||
name="高级会员月卡",
|
||||
description=None,
|
||||
monthly_price=39,
|
||||
yearly_price=0,
|
||||
quotas=None,
|
||||
display_order=2,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
assert len(plans) == 1
|
||||
p = plans[0]
|
||||
assert p["billing_cycle"] == "monthly"
|
||||
assert p["price_cents"] == 3900
|
||||
assert p["monthly_price_cents"] == 3900
|
||||
assert p["duration_days"] == 30
|
||||
assert p["features"] == {}
|
||||
|
||||
def test_both_cycles_expanded(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
plan_key="premium",
|
||||
name="高级会员",
|
||||
description=None,
|
||||
monthly_price=39,
|
||||
yearly_price=399,
|
||||
quotas={},
|
||||
display_order=1,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
plans = admin_catalog.get_membership_plans()
|
||||
|
||||
cycles = {p["billing_cycle"] for p in plans}
|
||||
assert cycles == {"yearly", "monthly"}
|
||||
|
||||
def test_no_session_returns_empty(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", None, create=True):
|
||||
assert admin_catalog.get_membership_plans() == []
|
||||
|
||||
|
||||
class TestPointsPackages:
|
||||
def test_package_mapping_with_bonus(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
package_key="pkg_100",
|
||||
name="100元充值包",
|
||||
price=100,
|
||||
credits=1000,
|
||||
bonus_credits=100,
|
||||
is_recommended=True,
|
||||
description="推荐",
|
||||
sort_order=4,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
packages = admin_catalog.get_points_packages()
|
||||
|
||||
assert len(packages) == 1
|
||||
pkg = packages[0]
|
||||
assert pkg["code"] == "pkg_100"
|
||||
assert pkg["points"] == 1100
|
||||
assert pkg["price_cents"] == 10000
|
||||
assert pkg["is_recommended"] is True
|
||||
assert pkg["unit_price"] == "¥0.091/积分"
|
||||
|
||||
def test_zero_credits_unit_price_safe(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
row = _row(
|
||||
package_key="pkg_0",
|
||||
name="空包",
|
||||
price=0,
|
||||
credits=0,
|
||||
bonus_credits=0,
|
||||
is_recommended=False,
|
||||
description=None,
|
||||
sort_order=0,
|
||||
)
|
||||
session = MagicMock()
|
||||
session.execute.return_value.fetchall.return_value = [row]
|
||||
sl = MagicMock(return_value=session)
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", sl, create=True):
|
||||
packages = admin_catalog.get_points_packages()
|
||||
|
||||
assert packages[0]["points"] == 0
|
||||
assert packages[0]["price_cents"] == 0
|
||||
assert packages[0]["unit_price"] == "¥0.000/积分"
|
||||
|
||||
def test_no_session_returns_empty(self):
|
||||
from packages.application.catalog import admin_catalog
|
||||
|
||||
with patch("packages.adapters.sqlalchemy_impl.session.SessionLocal", None, create=True):
|
||||
assert admin_catalog.get_points_packages() == []
|
||||
|
||||
|
||||
class TestPackagesRoute:
|
||||
def test_get_packages_route_returns_items(self):
|
||||
from app.api.routes.points import get_packages
|
||||
|
||||
cu = MagicMock()
|
||||
cu.user.member_type = None
|
||||
rows = [
|
||||
{
|
||||
"code": "pkg_10",
|
||||
"name": "10元充值包",
|
||||
"points": 100,
|
||||
"price_cents": 1000,
|
||||
"unit_price": "¥0.100/积分",
|
||||
}
|
||||
]
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_points_packages",
|
||||
return_value=rows,
|
||||
):
|
||||
resp = get_packages(current_user=cu)
|
||||
|
||||
assert len(resp.packages) == 1
|
||||
item = resp.packages[0]
|
||||
assert item.code == "pkg_10"
|
||||
assert item.points == 100
|
||||
assert item.price_cents == 1000
|
||||
@@ -95,29 +95,25 @@ class TestCheckEndpointWhenDisabled:
|
||||
# 不再走免费额度判定
|
||||
svc.check_daily_free_clip.assert_not_called()
|
||||
|
||||
def test_unknown_scene_allowed_when_disabled(self):
|
||||
"""任意 scene_key(含未知/已下线)系统关闭时都返回 allowed=True, cost=0。"""
|
||||
def test_unknown_scene_still_400_when_disabled(self):
|
||||
"""未知 scene 即使系统关闭也返回 400(参数校验先于开关)。"""
|
||||
from app.api.routes.points import check_points
|
||||
from app.schemas.points import PointsCheckRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
svc = MagicMock()
|
||||
svc.get_or_create_account.return_value = {"balance": 0}
|
||||
with (
|
||||
patch("app.api.routes.points._credits_enabled", return_value=False),
|
||||
patch("app.api.routes.points._get_service", return_value=svc),
|
||||
):
|
||||
resp = check_points(body=PointsCheckRequest(scene_key="nope"), current_user=_make_cu(), db=MagicMock())
|
||||
assert resp.allowed is True
|
||||
assert resp.required_points == 0
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
check_points(body=PointsCheckRequest(scene_key="nope"), current_user=_make_cu(), db=MagicMock())
|
||||
assert exc.value.status_code == 400
|
||||
|
||||
def test_check_enabled_calculates_cost(self):
|
||||
"""开关开启时保持原有计费校验(voice_clone_synth 正常计费)。"""
|
||||
"""开关开启时保持原有计费校验。"""
|
||||
from app.api.routes.points import check_points
|
||||
from app.schemas.points import PointsCheckRequest
|
||||
|
||||
svc = MagicMock()
|
||||
svc.check_daily_free_clip.return_value = False
|
||||
svc.get_or_create_account.return_value = {"balance": 100}
|
||||
body = PointsCheckRequest(scene_key="voice_clone_synth", quantity=1, duration_minutes=1)
|
||||
body = PointsCheckRequest(scene_key="ai_title", quantity=1)
|
||||
|
||||
with (
|
||||
patch("app.api.routes.points._credits_enabled", return_value=True),
|
||||
@@ -127,23 +123,6 @@ class TestCheckEndpointWhenDisabled:
|
||||
|
||||
assert resp.required_points == 2 # 免费用户 ceil(1*1.15)=2
|
||||
|
||||
def test_retired_scene_free_when_enabled(self):
|
||||
"""开关开启时,已下线场景返回 cost=0,直接放行。"""
|
||||
from app.api.routes.points import check_points
|
||||
from app.schemas.points import PointsCheckRequest
|
||||
|
||||
svc = MagicMock()
|
||||
svc.get_or_create_account.return_value = {"balance": 0}
|
||||
with (
|
||||
patch("app.api.routes.points._credits_enabled", return_value=True),
|
||||
patch("app.api.routes.points._get_service", return_value=svc),
|
||||
):
|
||||
for scene in ["ai_voice", "ai_title", "ai_video", "ai_digital_human", "nope"]:
|
||||
body = PointsCheckRequest(scene_key=scene, quantity=1)
|
||||
resp = check_points(body=body, current_user=_make_cu(), db=MagicMock())
|
||||
assert resp.required_points == 0, f"{scene} should be free"
|
||||
assert resp.allowed is True
|
||||
|
||||
|
||||
# ── /points/deduct:关闭时 no-op,余额不变 ────────────────────────────────
|
||||
|
||||
@@ -238,24 +217,13 @@ class TestQueryEndpointsRemainAvailable:
|
||||
|
||||
|
||||
class TestBusinessRoutesBypassWhenDisabled:
|
||||
def test_lipsync_route_has_no_points_logic(self):
|
||||
"""lipsync 路由:已移除手动扣点代码(不导入 PointsService/calculate_points_cost)。"""
|
||||
import inspect
|
||||
|
||||
def test_lipsync_route_skips_points(self):
|
||||
"""lipsync 创建任务路由:settings.points_enabled=False 时不构造 PointsService。"""
|
||||
from app.api.routes import lipsync as lipsync_mod
|
||||
|
||||
src = inspect.getsource(lipsync_mod)
|
||||
assert "PointsService" not in src
|
||||
assert "calculate_points_cost" not in src
|
||||
assert "_points_deducted" not in src
|
||||
|
||||
def test_tts_route_has_no_points_logic(self):
|
||||
"""tts 路由:已移除手动扣点代码(不导入 PointsService/calculate_points_cost)。"""
|
||||
import inspect
|
||||
assert bool(getattr(lipsync_mod.settings, "points_enabled", False)) is False
|
||||
|
||||
def test_tts_route_skips_points(self):
|
||||
from app.api.routes import tts as tts_mod
|
||||
|
||||
src = inspect.getsource(tts_mod)
|
||||
assert "PointsService" not in src
|
||||
assert "calculate_points_cost" not in src
|
||||
assert "_points_deducted" not in src
|
||||
assert bool(getattr(tts_mod.settings, "points_enabled", False)) is False
|
||||
|
||||
@@ -1,178 +0,0 @@
|
||||
"""tests for packages/shared/dashscope_client.py (#2159 Wan 3.0 DashScope client)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
_SINGLETON = "_DASHSCOPE_CLIENT_SINGLETON"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_singleton():
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
# 兼容实际 singleton 名
|
||||
for name in ("_DASHSCOPE_CLIENT_SINGLETON", "_dashscope_client"):
|
||||
if hasattr(d, name):
|
||||
setattr(d, name, None)
|
||||
yield
|
||||
for name in ("_DASHSCOPE_CLIENT_SINGLETON", "_dashscope_client"):
|
||||
if hasattr(d, name):
|
||||
setattr(d, name, None)
|
||||
|
||||
|
||||
def _make_settings(api_key="test-key"):
|
||||
return MagicMock(
|
||||
dashscope_api_key=api_key,
|
||||
dashscope_base_url="https://dashscope.aliyuncs.com/api/v1",
|
||||
dashscope_video_timeout=10,
|
||||
dashscope_video_poll_interval=0,
|
||||
video_dir="/tmp/videos",
|
||||
)
|
||||
|
||||
|
||||
class TestDashScopeAvailability:
|
||||
def test_unavailable_without_key(self):
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(api_key="")
|
||||
assert get_dashscope_client() is None
|
||||
|
||||
def test_available_with_key(self):
|
||||
from packages.shared.dashscope_client import get_dashscope_client
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = get_dashscope_client()
|
||||
assert c is not None
|
||||
assert c.is_available is True
|
||||
|
||||
|
||||
def _mock_stream_response(min_size=2048):
|
||||
"""构造 httpx.stream 上下文返回值,模拟返回若干字节的 mp4 内容。"""
|
||||
m = MagicMock()
|
||||
m.status_code = 200
|
||||
chunk = b"x" * min_size
|
||||
m.iter_bytes.return_value = [chunk]
|
||||
ctx = MagicMock()
|
||||
ctx.__enter__.return_value = m
|
||||
return ctx
|
||||
|
||||
|
||||
class TestDashScopeVideoGeneration:
|
||||
def test_happy_path_returns_video_path(self):
|
||||
"""POST create → GET poll (SUCCEEDED) → download → returns path + correct payload."""
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = d.DashScopeClient()
|
||||
|
||||
create_resp = MagicMock(status_code=200)
|
||||
create_resp.json.return_value = {"output": {"task_id": "task-abc"}}
|
||||
poll_resp = MagicMock(status_code=200)
|
||||
poll_resp.json.return_value = {
|
||||
"output": {"task_status": "SUCCEEDED", "video_url": "http://x/y.mp4"},
|
||||
"usage": {"billed_duration": 10},
|
||||
}
|
||||
# fake file: write enough bytes to pass the size>=1024 check
|
||||
m_open = mock_open()
|
||||
m_open.return_value.write.return_value = None
|
||||
fake_size = {"/tmp/videos/wan_task-abc.mp4": 4096}
|
||||
|
||||
def fake_getsize(p):
|
||||
return fake_size.get(p, 0)
|
||||
|
||||
def fake_exists(p):
|
||||
return p in fake_size
|
||||
|
||||
with (
|
||||
patch.object(d.httpx, "post", return_value=create_resp) as mock_post,
|
||||
patch.object(d.httpx, "get", return_value=poll_resp),
|
||||
patch.object(d.httpx, "stream", return_value=_mock_stream_response()),
|
||||
patch("packages.shared.dashscope_client.time.sleep"),
|
||||
patch("packages.shared.dashscope_client.os.makedirs"),
|
||||
patch("builtins.open", m_open),
|
||||
patch("packages.shared.dashscope_client.os.path.getsize", side_effect=fake_getsize),
|
||||
patch("packages.shared.dashscope_client.os.path.exists", side_effect=fake_exists),
|
||||
):
|
||||
res = c.video_generation(
|
||||
prompt="test",
|
||||
duration=5,
|
||||
ratio="9:16",
|
||||
resolution="720p",
|
||||
output_dir="/tmp/videos",
|
||||
)
|
||||
assert res is not None, "expected success"
|
||||
assert res["video_path"] == "/tmp/videos/wan_task-abc.mp4"
|
||||
_, kwargs = mock_post.call_args
|
||||
body = kwargs["json"]
|
||||
assert body["parameters"]["resolution"] == "720P"
|
||||
assert body["model"] == "wan3.0-video"
|
||||
|
||||
def test_create_http_error_returns_none(self):
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = d.DashScopeClient()
|
||||
err_resp = MagicMock(status_code=400, text="bad")
|
||||
err_resp.raise_for_status.side_effect = RuntimeError("bad")
|
||||
with patch.object(d.httpx, "post", return_value=err_resp):
|
||||
res = c.video_generation(
|
||||
prompt="test", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos"
|
||||
)
|
||||
assert res is None
|
||||
|
||||
def test_poll_failed_returns_none(self):
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = d.DashScopeClient()
|
||||
create_resp = MagicMock(status_code=200)
|
||||
create_resp.json.return_value = {"output": {"task_id": "task-abc"}}
|
||||
poll_resp = MagicMock(status_code=200)
|
||||
poll_resp.json.return_value = {"output": {"task_status": "FAILED", "message": "nope"}}
|
||||
with (
|
||||
patch.object(d.httpx, "post", return_value=create_resp),
|
||||
patch.object(d.httpx, "get", return_value=poll_resp),
|
||||
patch("packages.shared.dashscope_client.time.sleep"),
|
||||
):
|
||||
res = c.video_generation(
|
||||
prompt="test", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos"
|
||||
)
|
||||
assert res is None
|
||||
|
||||
def test_empty_prompt_returns_none(self):
|
||||
import packages.shared.dashscope_client as d
|
||||
|
||||
with patch("packages.shared.dashscope_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = d.DashScopeClient()
|
||||
assert (
|
||||
c.video_generation(prompt=" ", duration=5, ratio="9:16", resolution="720p", output_dir="/tmp/videos")
|
||||
is None
|
||||
)
|
||||
|
||||
def test_create_400_sets_last_video_error(self, tmp_path):
|
||||
"""创建任务 HTTP 400 时应写 last_video_error。"""
|
||||
from packages.shared import dashscope_client as dc
|
||||
|
||||
dc._DASHSCOPE_CLIENT_SINGLETON = None
|
||||
with patch.dict("os.environ", {"DASHSCOPE_API_KEY": "test-key"}):
|
||||
c = dc.DashScopeClient()
|
||||
r = MagicMock()
|
||||
r.status_code = 401
|
||||
r.text = '{"code":"InvalidApiKey","message":"bad key"}'
|
||||
r.raise_for_status.side_effect = httpx.HTTPStatusError("auth", request=MagicMock(), response=r)
|
||||
with patch.object(dc.httpx, "post", return_value=r), patch.object(dc, "time"):
|
||||
out = c.video_generation("p", output_dir=str(tmp_path))
|
||||
assert out is None
|
||||
err = c.get_last_video_error()
|
||||
assert err["error_code"] == "auth_error"
|
||||
assert c.last_video_error is not None
|
||||
@@ -1,25 +1,28 @@
|
||||
"""封面生成 — v1.6.2 起免费,不扣积分"""
|
||||
"""AI封面生成 积分扣点单元测试 (#1895 P2 step 2.7)"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
class TestGenerationCoverFree:
|
||||
def test_ai_cover_returns_zero_cost(self):
|
||||
import pytest
|
||||
|
||||
import packages.middleware.points_gate as _pg_module
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _enable(monkeypatch):
|
||||
monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True)
|
||||
yield
|
||||
|
||||
|
||||
class TestGenerationCoverPoints:
|
||||
def test_ai_cover_cost(self):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
assert calculate_points_cost("ai_cover", is_member=False, quantity=1) == 0
|
||||
assert calculate_points_cost("ai_cover", is_member=True, quantity=10) == 0
|
||||
assert calculate_points_cost("ai_cover", is_member=False) == 2
|
||||
assert calculate_points_cost("ai_cover", is_member=True, member_type="yearly") >= 0
|
||||
|
||||
def test_no_points_gate_decorator(self):
|
||||
def test_decorator_attached(self):
|
||||
from app.api.routes.generation_cover import generate_cover
|
||||
|
||||
assert not hasattr(generate_cover, "__wrapped__")
|
||||
|
||||
def test_endpoint_has_no_points_logic(self):
|
||||
import inspect
|
||||
|
||||
from app.api.routes import generation_cover as mod
|
||||
|
||||
src = inspect.getsource(mod)
|
||||
assert "PointsService" not in src
|
||||
assert "deduct_points" not in src
|
||||
assert hasattr(generate_cover, "__wrapped__"), "missing @points_gate"
|
||||
|
||||
@@ -1,31 +1,54 @@
|
||||
"""视频预览生成 — v1.6.2 起免费,不扣积分"""
|
||||
"""视频预览生成 积分扣点单元测试 (#1895 P2 step 2.5)"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
class TestGenerationPreviewFree:
|
||||
def test_ai_video_returns_zero_cost(self):
|
||||
import packages.middleware.points_gate as _pg_module
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _enable(monkeypatch):
|
||||
monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True)
|
||||
yield
|
||||
|
||||
|
||||
class TestGenerationPreviewPoints:
|
||||
def test_ai_video_cost(self):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
assert calculate_points_cost("ai_video", is_member=False) == 0
|
||||
assert calculate_points_cost("ai_video", is_member=True, member_type="monthly", duration_minutes=10) == 0
|
||||
assert calculate_points_cost("ai_video", is_member=False) == 4
|
||||
assert calculate_points_cost("ai_video", is_member=True, member_type="monthly") == 2
|
||||
|
||||
def test_no_points_gate_decorator(self):
|
||||
"""预览生成路由已移除 @points_gate。"""
|
||||
def test_insufficient_raises_402(self):
|
||||
from app.api.routes.generation_preview import create_preview_generation_task
|
||||
from app.schemas.generation_task import CreatePreviewGenerationTaskRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
db = MagicMock()
|
||||
cu = MagicMock()
|
||||
cu.user.id = "u1"
|
||||
cu.user.is_member = False
|
||||
cu.user.member_type = None
|
||||
req = CreatePreviewGenerationTaskRequest(template_id="t1", asset_ids=["a1"], preview_count=1)
|
||||
with patch("packages.domain.points_service.PointsService") as MS:
|
||||
svc = MagicMock()
|
||||
svc.check_daily_free_clip.return_value = False
|
||||
svc.deduct_points.return_value = {"success": False, "balance": 0}
|
||||
MS.return_value = svc
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
create_preview_generation_task(
|
||||
request=req,
|
||||
authenticated_user=cu,
|
||||
db=db,
|
||||
generation_task_repository=MagicMock(),
|
||||
asset_repo=MagicMock(),
|
||||
)
|
||||
assert ei.value.status_code == 402
|
||||
|
||||
def test_decorator_attached(self):
|
||||
from app.api.routes.generation_preview import create_preview_generation_task
|
||||
|
||||
# 移除装饰器后 __wrapped__ 不再存在
|
||||
assert not hasattr(create_preview_generation_task, "__wrapped__")
|
||||
|
||||
def test_endpoint_does_not_deduct_points(self):
|
||||
"""端点不再实例化 PointsService / 调用 deduct_points(直接走业务逻辑)。"""
|
||||
import inspect
|
||||
|
||||
from app.api.routes.generation_preview import create_preview_generation_task
|
||||
|
||||
src = inspect.getsource(create_preview_generation_task)
|
||||
assert "PointsService" not in src
|
||||
assert "deduct_points" not in src
|
||||
assert "calculate_points_cost" not in src
|
||||
assert hasattr(create_preview_generation_task, "__wrapped__"), "missing @points_gate"
|
||||
|
||||
@@ -1,28 +1,63 @@
|
||||
"""智能混剪任务 — v1.6.2 起免费,不扣积分"""
|
||||
"""视频生成 积分扣点单元测试 (#1895 P2 step 2.4)"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
import packages.middleware.points_gate as _pg_module
|
||||
|
||||
|
||||
class TestGenerationTasksFree:
|
||||
def test_ai_video_returns_zero_cost(self):
|
||||
@pytest.fixture(autouse=True)
|
||||
def _enable(monkeypatch):
|
||||
monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True)
|
||||
yield
|
||||
|
||||
|
||||
class TestGenerationTasksPoints:
|
||||
def test_ai_video_base_cost(self):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
assert calculate_points_cost("ai_video", is_member=False, duration_minutes=5) == 0
|
||||
assert calculate_points_cost("ai_video", is_member=True, duration_minutes=10) == 0
|
||||
assert calculate_points_cost("ai_video", is_member=False) == 4
|
||||
assert calculate_points_cost("ai_video", is_member=True, member_type="monthly") == 2
|
||||
|
||||
def test_no_points_gate_decorator(self):
|
||||
def test_ai_video_quantity_scales(self):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
c1 = calculate_points_cost("ai_video", is_member=False, quantity=1)
|
||||
c3 = calculate_points_cost("ai_video", is_member=False, quantity=3)
|
||||
assert c3 > c1
|
||||
|
||||
def test_insufficient_raises_402(self):
|
||||
from app.api.routes.generation_tasks import create_generation_task
|
||||
from app.schemas.generation_task import CreateGenerationTaskRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
db = MagicMock()
|
||||
cu = MagicMock()
|
||||
cu.user.id = "u1"
|
||||
cu.user.is_member = False
|
||||
cu.user.member_type = None
|
||||
req = CreateGenerationTaskRequest(template_id="t1", asset_ids=["a1"], count=1)
|
||||
with patch("packages.domain.points_service.PointsService") as MS:
|
||||
svc = MagicMock()
|
||||
svc.check_daily_free_clip.return_value = False
|
||||
svc.deduct_points.return_value = {"success": False, "balance": 0}
|
||||
MS.return_value = svc
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
create_generation_task(
|
||||
request=req,
|
||||
authenticated_user=cu,
|
||||
db=db,
|
||||
generation_task_repository=MagicMock(),
|
||||
project_repository=MagicMock(),
|
||||
asset_library_repository=MagicMock(),
|
||||
asset_repository=MagicMock(),
|
||||
)
|
||||
assert ei.value.status_code == 402
|
||||
|
||||
def test_decorator_attached(self):
|
||||
from app.api.routes.generation_tasks import create_generation_task
|
||||
|
||||
assert not hasattr(create_generation_task, "__wrapped__")
|
||||
|
||||
def test_create_task_accepts_request_without_points_block(self):
|
||||
"""路由函数签名不再做扣点,但参数 points_enabled/is_member/member_type 仍保留以兼容调用方。"""
|
||||
import inspect
|
||||
|
||||
from app.api.routes.generation_tasks import create_generation_task
|
||||
|
||||
sig = inspect.signature(create_generation_task)
|
||||
# 函数存在
|
||||
assert callable(create_generation_task)
|
||||
assert hasattr(create_generation_task, "__wrapped__"), "missing @points_gate"
|
||||
|
||||
@@ -1,376 +0,0 @@
|
||||
"""tests for packages/shared/jimeng_client.py (#2169 即梦 i2v 客户端)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import os
|
||||
from unittest.mock import MagicMock, mock_open, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_singleton():
|
||||
import packages.shared.jimeng_client as j
|
||||
|
||||
j._JIMENG_CLIENT_SINGLETON = None
|
||||
yield
|
||||
j._JIMENG_CLIENT_SINGLETON = None
|
||||
|
||||
|
||||
def _make_settings(ak="test-ak", sk="test-sk", req_key="jimeng_i2v_first_v30", timeout=60, poll_interval=2):
|
||||
return MagicMock(
|
||||
jimeng_ak=ak,
|
||||
jimeng_sk=sk,
|
||||
jimeng_base_url="https://visual.volcengineapi.com",
|
||||
jimeng_req_key=req_key,
|
||||
jimeng_video_timeout=timeout,
|
||||
jimeng_video_poll_interval=poll_interval,
|
||||
)
|
||||
|
||||
|
||||
# ── V4 签名单元测试 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestV4Signature:
|
||||
def test_sign_returns_required_headers(self):
|
||||
from packages.shared.jimeng_client import _sign_v4
|
||||
|
||||
headers = _sign_v4(
|
||||
ak="AK_TEST",
|
||||
sk="SK_TEST",
|
||||
method="POST",
|
||||
host="visual.volcengineapi.com",
|
||||
query={"Action": "CVSync2AsyncSubmitTask", "Version": "2022-08-31"},
|
||||
body_bytes=b'{"req_key":"jimeng_i2v_first_v30"}',
|
||||
x_date="20260101T120000Z",
|
||||
)
|
||||
assert headers["Content-Type"] == "application/json"
|
||||
assert headers["Host"] == "visual.volcengineapi.com"
|
||||
assert headers["X-Date"] == "20260101T120000Z"
|
||||
assert "X-Content-Sha256" in headers
|
||||
assert headers["Authorization"].startswith("HMAC-SHA256 Credential=AK_TEST/20260101/cn-north-1/cv/request")
|
||||
assert "SignedHeaders=content-type;host;x-content-sha256;x-date" in headers["Authorization"]
|
||||
assert "Signature=" in headers["Authorization"]
|
||||
# 签名是 64 字符 hex
|
||||
sig = headers["Authorization"].split("Signature=")[-1]
|
||||
assert len(sig) == 64
|
||||
assert all(c in "0123456789abcdef" for c in sig)
|
||||
|
||||
def test_sign_deterministic(self):
|
||||
"""相同输入必须产生相同签名(幂等)。"""
|
||||
from packages.shared.jimeng_client import _sign_v4
|
||||
|
||||
kwargs = dict(
|
||||
ak="AK",
|
||||
sk="SK",
|
||||
method="POST",
|
||||
host="h",
|
||||
query={"A": "1", "B": "2"},
|
||||
body_bytes=b"{}",
|
||||
x_date="20260101T000000Z",
|
||||
)
|
||||
h1 = _sign_v4(**kwargs)
|
||||
h2 = _sign_v4(**kwargs)
|
||||
assert h1["Authorization"] == h2["Authorization"]
|
||||
assert h1["X-Content-Sha256"] == h2["X-Content-Sha256"]
|
||||
|
||||
def test_sign_different_body_different_sig(self):
|
||||
from packages.shared.jimeng_client import _sign_v4
|
||||
|
||||
base = dict(ak="AK", sk="SK", method="POST", host="h", query={}, x_date="20260101T000000Z")
|
||||
h1 = _sign_v4(body_bytes=b"a", **base)
|
||||
h2 = _sign_v4(body_bytes=b"b", **base)
|
||||
assert h1["Authorization"] != h2["Authorization"]
|
||||
|
||||
def test_payload_sha256_matches(self):
|
||||
import hashlib
|
||||
|
||||
from packages.shared.jimeng_client import _sign_v4
|
||||
|
||||
body = b'{"prompt":"hello"}'
|
||||
h = _sign_v4("ak", "sk", "POST", "h", {}, body, "20260101T000000Z")
|
||||
expected = hashlib.sha256(body).hexdigest()
|
||||
assert h["X-Content-Sha256"] == expected
|
||||
|
||||
def test_norm_query_sorted_and_encoded(self):
|
||||
from packages.shared.jimeng_client import _norm_query
|
||||
|
||||
q = _norm_query({"B": "2", "A": "1", "C": "a b"})
|
||||
# key 排序 + 空格→%20
|
||||
assert q.startswith("A=1")
|
||||
assert "B=2" in q
|
||||
assert "C=a%20b" in q
|
||||
|
||||
|
||||
# ── 可用性 / 单例 ─────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAvailability:
|
||||
def test_unavailable_without_ak_sk(self):
|
||||
from packages.shared.jimeng_client import JimengClient, get_jimeng_client
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(ak="", sk="")
|
||||
# 重置单例
|
||||
import packages.shared.jimeng_client as j
|
||||
|
||||
j._JIMENG_CLIENT_SINGLETON = None
|
||||
assert get_jimeng_client() is None
|
||||
|
||||
def test_available_with_ak_sk(self):
|
||||
from packages.shared.jimeng_client import get_jimeng_client
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
import packages.shared.jimeng_client as j
|
||||
|
||||
j._JIMENG_CLIENT_SINGLETON = None
|
||||
c = get_jimeng_client()
|
||||
assert c is not None
|
||||
assert c.is_available is True
|
||||
assert c.req_key == "jimeng_i2v_first_v30"
|
||||
|
||||
|
||||
# ── 错误分类 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestClassifyError:
|
||||
def test_50411_is_portrait_intercept_non_retryable(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, msg, retry = _classify_jimeng_error(200, '{"code":50411,"message":"Pre Img Risk"}', 50411)
|
||||
assert code == "portrait_intercept"
|
||||
assert retry is False
|
||||
|
||||
def test_50429_is_rate_limit_retryable(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, msg, retry = _classify_jimeng_error(200, "", 50429)
|
||||
assert code == "rate_limit"
|
||||
assert retry is True
|
||||
|
||||
def test_50430_is_rate_limit(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, _, _ = _classify_jimeng_error(200, "", 50430)
|
||||
assert code == "rate_limit"
|
||||
|
||||
def test_50500_is_network_error_retryable(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, _, retry = _classify_jimeng_error(200, "", 50500)
|
||||
assert code == "network_error"
|
||||
assert retry is True
|
||||
|
||||
def test_50412_is_invalid_param_non_retryable(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, _, retry = _classify_jimeng_error(200, "", 50412)
|
||||
assert code == "invalid_param"
|
||||
assert retry is False
|
||||
|
||||
def test_401_auth(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, msg, retry = _classify_jimeng_error(401, "auth fail", None)
|
||||
assert code == "auth_error"
|
||||
assert retry is False
|
||||
|
||||
def test_400_text_audit(self):
|
||||
from packages.shared.jimeng_client import _classify_jimeng_error
|
||||
|
||||
code, _, _ = _classify_jimeng_error(400, "text error", None)
|
||||
assert code == "invalid_param"
|
||||
|
||||
|
||||
# ── video_generation 主流程 ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVideoGenerationHappyPath:
|
||||
def test_missing_ak_returns_none(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(ak="", sk="")
|
||||
c = JimengClient()
|
||||
assert c.video_generation("hi", image_url="http://x/y.jpg") is None
|
||||
err = c.last_video_error
|
||||
assert err["error_code"] == "auth_error"
|
||||
|
||||
def test_empty_prompt_returns_none(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = JimengClient()
|
||||
assert c.video_generation(" ", image_url="http://x/y.jpg") is None
|
||||
assert c.last_video_error["error_code"] == "invalid_param"
|
||||
|
||||
def test_empty_image_url_returns_none(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = JimengClient()
|
||||
assert c.video_generation("prompt", image_url="") is None
|
||||
assert c.last_video_error["error_code"] == "invalid_param"
|
||||
|
||||
def test_duration_5s_frames_121(self):
|
||||
"""5s → frames=121,10s→frames=241,>10s 截断到10s。"""
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=1, poll_interval=0)
|
||||
c = JimengClient()
|
||||
|
||||
captured = {}
|
||||
|
||||
def fake_submit(prompt, image_url, frames, seed=-1):
|
||||
captured["frames"] = frames
|
||||
return "task-xyz"
|
||||
|
||||
def fake_poll(tid):
|
||||
captured["tid"] = tid
|
||||
return "http://example.com/v.mp4"
|
||||
|
||||
def fake_download(url, out_dir, tid):
|
||||
captured["url"] = url
|
||||
return "/tmp/fake.mp4"
|
||||
|
||||
# 构造一个假文件
|
||||
os.makedirs("/tmp", exist_ok=True)
|
||||
with open("/tmp/fake.mp4", "wb") as f:
|
||||
f.write(b"x" * 2048)
|
||||
|
||||
with (
|
||||
patch.object(c, "_submit_task", side_effect=fake_submit),
|
||||
patch.object(c, "_poll_result", side_effect=fake_poll),
|
||||
patch.object(c, "_download_video", side_effect=fake_download),
|
||||
):
|
||||
r = c.video_generation("test", image_url="http://x/y.jpg", duration=5, output_dir="/tmp")
|
||||
assert r is not None
|
||||
assert captured["frames"] == 121
|
||||
assert r["usage"]["duration_seconds"] == 5
|
||||
assert r["usage"]["billing_mode"] == "per_second"
|
||||
assert r["usage"]["req_key"] == "jimeng_i2v_first_v30"
|
||||
|
||||
def test_duration_10s_frames_241(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=1, poll_interval=0)
|
||||
c = JimengClient()
|
||||
captured = {}
|
||||
|
||||
def fake_submit(prompt, image_url, frames, seed=-1):
|
||||
captured["frames"] = frames
|
||||
return "tid"
|
||||
|
||||
def fake_poll(tid):
|
||||
return "http://x/v.mp4"
|
||||
|
||||
def fake_download(url, out_dir, tid):
|
||||
with open("/tmp/fake2.mp4", "wb") as f:
|
||||
f.write(b"x" * 2048)
|
||||
return "/tmp/fake2.mp4"
|
||||
|
||||
with (
|
||||
patch.object(c, "_submit_task", side_effect=fake_submit),
|
||||
patch.object(c, "_poll_result", side_effect=fake_poll),
|
||||
patch.object(c, "_download_video", side_effect=fake_download),
|
||||
):
|
||||
r = c.video_generation("hi", image_url="http://x/y.jpg", duration=10, output_dir="/tmp")
|
||||
assert captured["frames"] == 241
|
||||
assert r["usage"]["duration_seconds"] == 10
|
||||
|
||||
def test_duration_over_10s_truncates_to_10s(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=1, poll_interval=0)
|
||||
c = JimengClient()
|
||||
captured = {}
|
||||
|
||||
def fake_submit(prompt, image_url, frames, seed=-1):
|
||||
captured["frames"] = frames
|
||||
return "tid"
|
||||
|
||||
def fake_poll(tid):
|
||||
return "http://x/v.mp4"
|
||||
|
||||
def fake_download(url, out_dir, tid):
|
||||
with open("/tmp/fake3.mp4", "wb") as f:
|
||||
f.write(b"x" * 2048)
|
||||
return "/tmp/fake3.mp4"
|
||||
|
||||
with (
|
||||
patch.object(c, "_submit_task", side_effect=fake_submit),
|
||||
patch.object(c, "_poll_result", side_effect=fake_poll),
|
||||
patch.object(c, "_download_video", side_effect=fake_download),
|
||||
):
|
||||
r = c.video_generation("hi", image_url="http://x/y.jpg", duration=30, output_dir="/tmp")
|
||||
assert captured["frames"] == 241
|
||||
assert r["usage"]["duration_seconds"] == 10
|
||||
|
||||
|
||||
class TestSubmitTaskErrors:
|
||||
def test_submit_50411_writes_portrait_intercept(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = JimengClient()
|
||||
fake_resp = MagicMock(status_code=200, text='{"code":50411,"message":"Pre Img Risk Not Pass"}')
|
||||
fake_resp.json.return_value = {"code": 50411, "message": "Pre Img Risk Not Pass"}
|
||||
with patch("packages.shared.jimeng_client.httpx.request", return_value=fake_resp):
|
||||
tid = c._submit_task("p", "http://x/y.jpg", frames=121)
|
||||
assert tid is None
|
||||
assert c.last_video_error["error_code"] == "portrait_intercept"
|
||||
|
||||
def test_submit_returns_task_id(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings()
|
||||
c = JimengClient()
|
||||
fake_resp = MagicMock(status_code=200, text='{"code":10000,"data":{"task_id":"abc"}}')
|
||||
fake_resp.json.return_value = {"code": 10000, "data": {"task_id": "abc"}}
|
||||
with patch("packages.shared.jimeng_client.httpx.request", return_value=fake_resp):
|
||||
tid = c._submit_task("p", "http://x/y.jpg", frames=121)
|
||||
assert tid == "abc"
|
||||
|
||||
|
||||
class TestPollResult:
|
||||
def test_poll_done_returns_video_url(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=10, poll_interval=0)
|
||||
c = JimengClient()
|
||||
done_resp = MagicMock(status_code=200)
|
||||
done_resp.json.return_value = {"code": 10000, "data": {"status": "done", "video_url": "http://x/v.mp4"}}
|
||||
with (
|
||||
patch("packages.shared.jimeng_client.httpx.request", return_value=done_resp),
|
||||
patch("packages.shared.jimeng_client.time.sleep"),
|
||||
):
|
||||
url = c._poll_result("abc")
|
||||
assert url == "http://x/v.mp4"
|
||||
|
||||
def test_poll_timeout_returns_none(self):
|
||||
from packages.shared.jimeng_client import JimengClient
|
||||
|
||||
with patch("packages.shared.jimeng_client.get_shared_settings") as ms:
|
||||
ms.return_value = _make_settings(timeout=1, poll_interval=0)
|
||||
c = JimengClient()
|
||||
queue_resp = MagicMock(status_code=200)
|
||||
queue_resp.json.return_value = {"code": 10000, "data": {"status": "in_queue"}}
|
||||
# time.time 会被调用,模拟超时
|
||||
with (
|
||||
patch("packages.shared.jimeng_client.httpx.request", return_value=queue_resp),
|
||||
patch("packages.shared.jimeng_client.time.sleep"),
|
||||
):
|
||||
url = c._poll_result("abc")
|
||||
assert url is None
|
||||
assert c.last_video_error["error_code"] == "network_error"
|
||||
assert "超时" in c.last_video_error["user_message"]
|
||||
@@ -1,16 +1,15 @@
|
||||
"""lipsync 口型同步 — v1.6.2 起免费,不扣积分"""
|
||||
"""lipsync 积分扣点单元测试 (#1895 P2 step 2.2)"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
|
||||
def _cu(user_id="u1", is_member=False, member_type=None):
|
||||
def _make_cu(user_id="user-1", is_member=False, member_type=None):
|
||||
cu = MagicMock()
|
||||
cu.user.id = user_id
|
||||
cu.user.is_member = is_member
|
||||
@@ -18,6 +17,99 @@ def _cu(user_id="u1", is_member=False, member_type=None):
|
||||
return cu
|
||||
|
||||
|
||||
class TestLipsyncDurationEstimate:
|
||||
@pytest.mark.parametrize(
|
||||
"text,expected",
|
||||
[
|
||||
("你好", 1.0),
|
||||
("你" * 240, 1.0),
|
||||
("你" * 241, 2.0),
|
||||
("你" * 1000, 5.0),
|
||||
],
|
||||
)
|
||||
def test_text_estimate(self, text, expected):
|
||||
est = max(1.0, math.ceil(len(text) / 240))
|
||||
assert est == expected
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"seconds,expected",
|
||||
[
|
||||
(30, 1.0),
|
||||
(60, 1.0),
|
||||
(61, 2.0),
|
||||
(120, 2.0),
|
||||
(180, 3.0),
|
||||
],
|
||||
)
|
||||
def test_audio_duration_estimate(self, seconds, expected):
|
||||
est = max(1.0, math.ceil(seconds / 60.0))
|
||||
assert est == expected
|
||||
|
||||
|
||||
class TestLipsyncPointsDeduction:
|
||||
def _deduct(self, text="你好", audio_duration=None, enabled=True, success=True, balance=100, **cu_kw):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
svc = MagicMock() if enabled else None
|
||||
cu = _make_cu(**cu_kw)
|
||||
if svc is None:
|
||||
return 0, cu
|
||||
if audio_duration and audio_duration > 0:
|
||||
est = max(1.0, math.ceil(audio_duration / 60.0))
|
||||
elif text:
|
||||
est = max(1.0, math.ceil(len(text) / 240))
|
||||
else:
|
||||
est = 1.0
|
||||
cost = calculate_points_cost(
|
||||
"ai_digital_human",
|
||||
is_member=getattr(cu.user, "is_member", False),
|
||||
duration_minutes=est,
|
||||
member_type=getattr(cu.user, "member_type", None),
|
||||
)
|
||||
svc.deduct_points.return_value = {"success": success, "balance": balance}
|
||||
res = svc.deduct_points(cu.user.id, cost, "ai_digital_human", MagicMock())
|
||||
if not res["success"]:
|
||||
raise HTTPException(status_code=402, detail={"code": "INSUFFICIENT_POINTS"})
|
||||
return cost, cu
|
||||
|
||||
def test_disabled(self):
|
||||
cost, _ = self._deduct(enabled=False)
|
||||
assert cost == 0
|
||||
|
||||
def test_short_text_min_1min(self):
|
||||
cost, _ = self._deduct(text="你好")
|
||||
assert cost >= 15 # 15 base/min for free user × 1.15
|
||||
|
||||
def test_audio_duration_used(self):
|
||||
cost_long, _ = self._deduct(audio_duration=180) # 3min
|
||||
cost_short, _ = self._deduct(audio_duration=30) # 1min
|
||||
assert cost_long > cost_short
|
||||
|
||||
def test_insufficient_402(self):
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
self._deduct(text="你" * 500, success=False, balance=0)
|
||||
assert ei.value.status_code == 402
|
||||
|
||||
def test_member_cheaper(self):
|
||||
cm, _ = self._deduct(text="你" * 500, is_member=True, member_type="yearly")
|
||||
cf, _ = self._deduct(text="你" * 500, is_member=False)
|
||||
assert cm < cf
|
||||
|
||||
|
||||
# ── 直接调用 create_lipsync_job 覆盖扣点/402/退费分支 ──
|
||||
import importlib
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import patch
|
||||
|
||||
import packages.middleware.points_gate as _pg_module
|
||||
|
||||
|
||||
# Ensure the enable-gate fixture for lipsync also covers @points_gate (if any)
|
||||
# (the existing autouse _enable is below; importlib to avoid duplicate)
|
||||
def _do_enable(monkeypatch):
|
||||
monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True)
|
||||
|
||||
|
||||
def _body(**kw):
|
||||
b = MagicMock()
|
||||
defaults = dict(
|
||||
@@ -38,76 +130,113 @@ def _body(**kw):
|
||||
return b
|
||||
|
||||
|
||||
class TestLipsyncFree:
|
||||
"""lipsync 已移除手动扣点,业务异常仍按原状态码抛出。"""
|
||||
def _cu(user_id="u1", is_member=False, member_type=None):
|
||||
cu = MagicMock()
|
||||
cu.user.id = user_id
|
||||
cu.user.is_member = is_member
|
||||
cu.user.member_type = member_type
|
||||
return cu
|
||||
|
||||
def test_ai_digital_human_returns_zero_cost(self):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
assert calculate_points_cost("ai_digital_human", is_member=False, duration_minutes=1) == 0
|
||||
assert calculate_points_cost("ai_digital_human", is_member=True, duration_minutes=10) == 0
|
||||
|
||||
def test_module_has_no_points_imports(self):
|
||||
import inspect
|
||||
|
||||
from app.api.routes import lipsync as mod
|
||||
|
||||
src = inspect.getsource(mod)
|
||||
assert "PointsService" not in src
|
||||
assert "calculate_points_cost" not in src
|
||||
assert "_points_deducted" not in src
|
||||
assert "settings" not in src # settings was only used for points_enabled
|
||||
|
||||
def test_docstring_at_top_of_create_lipsync_job(self):
|
||||
"""扣点块删除后,docstring 必须在函数体第一行(防止函数体中段 docstring 丢失)。"""
|
||||
import ast
|
||||
import inspect
|
||||
|
||||
class TestLipsyncEndpointPoints:
|
||||
def test_insufficient_raises_402(self, monkeypatch):
|
||||
_do_enable(monkeypatch)
|
||||
from app.api.routes.lipsync import create_lipsync_job
|
||||
|
||||
src = inspect.getsource(create_lipsync_job)
|
||||
tree = ast.parse(src)
|
||||
fn = tree.body[0]
|
||||
# docstring 应为函数体第一条语句
|
||||
assert (
|
||||
isinstance(fn.body[0], ast.Expr)
|
||||
and isinstance(fn.body[0].value, ast.Constant)
|
||||
and isinstance(fn.body[0].value.value, str)
|
||||
), "create_lipsync_job docstring 不在函数体开头"
|
||||
db = MagicMock()
|
||||
svc = MagicMock()
|
||||
ps = MagicMock()
|
||||
ps.deduct_points.return_value = {"success": False, "balance": 0}
|
||||
fs = MagicMock(points_enabled=True)
|
||||
with (
|
||||
patch("app.api.routes.lipsync.PointsService", return_value=ps),
|
||||
patch("app.api.routes.lipsync.settings", fs),
|
||||
):
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
create_lipsync_job(body=_body(script_text="你" * 500), current_user=_cu(), db=db, svc=svc)
|
||||
assert ei.value.status_code == 402
|
||||
|
||||
def test_docstring_at_top_of_preview_tts(self):
|
||||
import ast
|
||||
import inspect
|
||||
|
||||
from app.api.routes.lipsync import preview_tts
|
||||
|
||||
src = inspect.getsource(preview_tts)
|
||||
tree = ast.parse(src)
|
||||
fn = tree.body[0]
|
||||
assert (
|
||||
isinstance(fn.body[0], ast.Expr)
|
||||
and isinstance(fn.body[0].value, ast.Constant)
|
||||
and isinstance(fn.body[0].value.value, str)
|
||||
), "preview_tts docstring 不在函数体开头"
|
||||
|
||||
def test_value_error_still_raises_400(self):
|
||||
"""业务异常仍抛 400(不再退费)。"""
|
||||
def test_value_error_refunds(self, monkeypatch):
|
||||
_do_enable(monkeypatch)
|
||||
from app.api.routes.lipsync import create_lipsync_job
|
||||
|
||||
db = MagicMock()
|
||||
svc = MagicMock()
|
||||
svc.create_job.side_effect = ValueError("bad input")
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc)
|
||||
assert ei.value.status_code == 400
|
||||
ps = MagicMock()
|
||||
ps.deduct_points.return_value = {"success": True, "balance": 99}
|
||||
fs = MagicMock(points_enabled=True)
|
||||
with (
|
||||
patch("app.api.routes.lipsync.PointsService", return_value=ps),
|
||||
patch("app.api.routes.lipsync.settings", fs),
|
||||
):
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc)
|
||||
assert ei.value.status_code == 400
|
||||
assert ps.refund_points.called
|
||||
|
||||
def test_success_returns_job(self):
|
||||
def test_mediakit_error_refunds(self, monkeypatch):
|
||||
_do_enable(monkeypatch)
|
||||
from app.api.routes.lipsync import create_lipsync_job
|
||||
from app.services.mediakit_client import MediaKitError
|
||||
|
||||
db = MagicMock()
|
||||
svc = MagicMock()
|
||||
svc.create_job.side_effect = MediaKitError("fail", code="InvalidInput")
|
||||
ps = MagicMock()
|
||||
ps.deduct_points.return_value = {"success": True, "balance": 99}
|
||||
fs = MagicMock(points_enabled=True)
|
||||
with (
|
||||
patch("app.api.routes.lipsync.PointsService", return_value=ps),
|
||||
patch("app.api.routes.lipsync.settings", fs),
|
||||
):
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc)
|
||||
assert ei.value.status_code == 400
|
||||
assert ps.refund_points.called
|
||||
|
||||
def test_generic_exception_refunds(self, monkeypatch):
|
||||
_do_enable(monkeypatch)
|
||||
from app.api.routes.lipsync import create_lipsync_job
|
||||
|
||||
db = MagicMock()
|
||||
svc = MagicMock()
|
||||
svc.create_job.side_effect = RuntimeError("boom")
|
||||
ps = MagicMock()
|
||||
ps.deduct_points.return_value = {"success": True, "balance": 99}
|
||||
fs = MagicMock(points_enabled=True)
|
||||
with (
|
||||
patch("app.api.routes.lipsync.PointsService", return_value=ps),
|
||||
patch("app.api.routes.lipsync.settings", fs),
|
||||
):
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc)
|
||||
assert ei.value.status_code == 400
|
||||
assert ps.refund_points.called
|
||||
|
||||
def test_audio_duration_estimation(self, monkeypatch):
|
||||
_do_enable(monkeypatch)
|
||||
from app.api.routes.lipsync import create_lipsync_job
|
||||
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
db = MagicMock()
|
||||
svc = MagicMock()
|
||||
job = SimpleNamespace(id="job-1", status="queued")
|
||||
svc.create_job.return_value = job
|
||||
# 不再依赖 settings/PointsService patch
|
||||
result = create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc)
|
||||
assert result is job
|
||||
ps = MagicMock()
|
||||
ps.deduct_points.return_value = {"success": True, "balance": 99}
|
||||
fs = MagicMock(points_enabled=True)
|
||||
with (
|
||||
patch("app.api.routes.lipsync.PointsService", return_value=ps),
|
||||
patch("app.api.routes.lipsync.settings", fs),
|
||||
):
|
||||
create_lipsync_job(
|
||||
body=_body(audio_url="http://x/a.mp3", audio_duration=180, script_text=None),
|
||||
current_user=_cu(),
|
||||
db=db,
|
||||
svc=svc,
|
||||
)
|
||||
# 180 seconds -> 3 minutes; assert deduct called with cost >= 15*3
|
||||
args = ps.deduct_points.call_args[0]
|
||||
assert args[1] >= calculate_points_cost("ai_digital_human", is_member=False, duration_minutes=3)
|
||||
|
||||
@@ -44,7 +44,7 @@ class TestExtractKwargs:
|
||||
|
||||
class TestPointsGateSync:
|
||||
def test_no_user_raises_401(self):
|
||||
@points_gate("voice_clone_synth")
|
||||
@points_gate("ai_rewrite")
|
||||
def my_func(db=None):
|
||||
return "ok"
|
||||
|
||||
@@ -53,7 +53,7 @@ class TestPointsGateSync:
|
||||
assert exc_info.value.status_code == 401
|
||||
|
||||
def test_no_db_raises_500(self):
|
||||
@points_gate("voice_clone_synth")
|
||||
@points_gate("ai_rewrite")
|
||||
def my_func(current_user=None, db=None):
|
||||
return "ok"
|
||||
|
||||
@@ -85,14 +85,7 @@ class TestPointsGateExecuteLogic:
|
||||
with patch("packages.domain.points_service.PointsService", return_value=mock_svc):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
_execute_with_gate(
|
||||
my_func,
|
||||
(),
|
||||
{"current_user": cu, "db": db},
|
||||
"voice_clone_synth",
|
||||
per_unit=10,
|
||||
unit_field=None,
|
||||
quantity_field=None,
|
||||
is_async=False,
|
||||
my_func, (), {"current_user": cu, "db": db}, "ai_rewrite", None, None, None, is_async=False
|
||||
)
|
||||
assert exc_info.value.status_code == 402
|
||||
|
||||
@@ -122,7 +115,7 @@ class TestPointsGateExecuteLogic:
|
||||
my_func,
|
||||
(),
|
||||
{"current_user": cu, "db": db},
|
||||
"voice_clone_synth",
|
||||
"ai_rewrite",
|
||||
per_unit=10,
|
||||
unit_field=None,
|
||||
quantity_field=None,
|
||||
@@ -146,7 +139,7 @@ class TestPointsGateExecuteLogic:
|
||||
failing_func,
|
||||
(),
|
||||
{"current_user": cu, "db": db},
|
||||
"voice_clone_synth",
|
||||
"ai_rewrite",
|
||||
per_unit=10,
|
||||
unit_field=None,
|
||||
quantity_field=None,
|
||||
@@ -154,21 +147,21 @@ class TestPointsGateExecuteLogic:
|
||||
)
|
||||
mock_svc.refund_points.assert_called_once()
|
||||
|
||||
def test_retired_scene_passes_through_with_zero_deduction(self):
|
||||
"""已下线场景(如 ai_video/ai_rewrite/ai_voice 等)直接放行,不扣积分。"""
|
||||
def test_ai_video_free_quota_for_free_user(self):
|
||||
cu = _make_current_user(is_member=False)
|
||||
db = MagicMock()
|
||||
mock_svc = MagicMock()
|
||||
mock_svc.check_daily_free_clip.return_value = True
|
||||
mock_svc.record_daily_free_clip.return_value = True
|
||||
|
||||
def my_func(current_user=cu, db=db, **kwargs):
|
||||
return kwargs.get("_points_deducted", -1)
|
||||
return kwargs.get("_is_free_quota", False)
|
||||
|
||||
# 不应调用 PointsService
|
||||
with patch("packages.domain.points_service.PointsService") as mock_svc_cls:
|
||||
with patch("packages.domain.points_service.PointsService", return_value=mock_svc):
|
||||
result = _execute_with_gate(
|
||||
my_func, (), {"current_user": cu, "db": db}, "ai_video", None, None, None, is_async=False
|
||||
)
|
||||
assert result == 0
|
||||
mock_svc_cls.assert_not_called()
|
||||
assert result is True
|
||||
|
||||
|
||||
class TestPointsGateAsync:
|
||||
@@ -179,7 +172,7 @@ class TestPointsGateAsync:
|
||||
mock_svc = MagicMock()
|
||||
mock_svc.deduct_points.return_value = {"success": True, "balance": 90, "transaction_id": "t1"}
|
||||
|
||||
@points_gate("voice_clone_synth", per_unit=5)
|
||||
@points_gate("ai_rewrite", per_unit=5)
|
||||
async def my_async_func(current_user=None, db=None, **kwargs):
|
||||
return kwargs.get("_points_deducted", 0)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
覆盖:
|
||||
- P0-1: POST /points/recharge 返回 pay_params / points_amount / expire_at
|
||||
- P0-2: POST /points/check 任意 scene_key 均可查询(已下线场景返回 cost=0,不报错)
|
||||
- P0-2: POST /points/check 未知 scene_key 返回 400(非 500)
|
||||
- P1-3: GET /points/rules 返回 description 字段
|
||||
- P1-6: GET /subscription/plans 返回档位列表
|
||||
- P1-7: multiplier 实际扣费一致(calculate_points_cost 统一应用)
|
||||
@@ -76,40 +76,39 @@ class TestRechargeOrderResponse:
|
||||
assert exc.value.status_code == 400
|
||||
|
||||
|
||||
# ── P0-2: check 任意 scene_key(已下线场景返回 cost=0) ──────────────────
|
||||
# ── P0-2: check unknown scene → 400 ───────────────────────────────────
|
||||
|
||||
|
||||
class TestCheckPointsUnknownScene:
|
||||
def test_unknown_scene_returns_zero_cost_not_error(self):
|
||||
"""任意 scene_key 均可查询,已下线/未知场景返回 cost=0(免费放行)。"""
|
||||
def test_unknown_scene_returns_400_not_500(self):
|
||||
"""未知 scene_key(如 ai_script)应返回 400 UNKNOWN_SCENE,而不是 500。"""
|
||||
from app.api.routes.points import check_points
|
||||
from app.schemas.points import PointsCheckRequest
|
||||
|
||||
svc = MagicMock()
|
||||
svc.get_or_create_account.return_value = {"balance": 0}
|
||||
db = MagicMock()
|
||||
cu = _make_cu()
|
||||
body = PointsCheckRequest(scene_key="ai_script", quantity=1)
|
||||
|
||||
with (
|
||||
patch("app.api.routes.points._credits_enabled", return_value=True),
|
||||
patch("app.api.routes.points._get_service", return_value=svc),
|
||||
):
|
||||
for scene in ["ai_script", "ai_voice", "ai_video", "ai_title", "ai_cover", "nonexistent"]:
|
||||
body = PointsCheckRequest(scene_key=scene, quantity=1)
|
||||
resp = check_points(body=body, current_user=cu, db=db)
|
||||
assert resp.required_points == 0, f"{scene} should be free"
|
||||
assert resp.allowed is True
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
check_points(body=body, current_user=cu, db=db)
|
||||
assert exc.value.status_code == 400
|
||||
detail = exc.value.detail
|
||||
assert detail["code"] == "UNKNOWN_SCENE"
|
||||
assert "ai_script" in detail["message"]
|
||||
assert "ai_voice" in detail["valid_scenes"]
|
||||
assert "ai_title" in detail["valid_scenes"]
|
||||
|
||||
def test_voice_clone_synth_still_charges(self):
|
||||
"""合法付费场景 voice_clone_synth 正常计费:免费用户 1 分钟 = ceil(1*1.15)=2 积分。"""
|
||||
def test_known_scene_still_works(self):
|
||||
"""合法 scene_key 正常返回,免费用户 ai_voice 1 分钟 = 2 积分。"""
|
||||
from app.api.routes.points import check_points
|
||||
from app.schemas.points import PointsCheckRequest
|
||||
|
||||
svc = MagicMock()
|
||||
svc.check_daily_free_clip.return_value = False
|
||||
svc.get_or_create_account.return_value = {"balance": 50}
|
||||
db = MagicMock()
|
||||
cu = _make_cu()
|
||||
body = PointsCheckRequest(scene_key="voice_clone_synth", quantity=1, duration_minutes=1)
|
||||
body = PointsCheckRequest(scene_key="ai_voice", quantity=1, duration_minutes=1)
|
||||
|
||||
with (
|
||||
patch("app.api.routes.points._credits_enabled", return_value=True),
|
||||
@@ -129,9 +128,7 @@ class TestPointsRulesDescription:
|
||||
from app.api.routes.points import get_rules
|
||||
|
||||
resp = get_rules(_current_user=_make_cu())
|
||||
# 场景列表包含 voice_clone_train / voice_clone_synth / viral_video(爆款视频为动态定价)
|
||||
keys = {r.scene_key for r in resp.rules}
|
||||
assert {"voice_clone_train", "voice_clone_synth", "viral_video"}.issubset(keys)
|
||||
assert len(resp.rules) >= 9
|
||||
for rule in resp.rules:
|
||||
assert rule.description, f"{rule.scene_key} missing description"
|
||||
assert isinstance(rule.description, str)
|
||||
@@ -173,63 +170,43 @@ class TestSubscriptionPlans:
|
||||
_spec.loader.exec_module(_mod)
|
||||
return _mod.list_membership_plans
|
||||
|
||||
def test_plans_endpoint_reads_admin_table(self):
|
||||
"""/subscription/plans 改读管理后台 plans 表:返回 catalog 服务提供的真实档位。"""
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
real_plan = {
|
||||
"plan_id": "premium_yearly",
|
||||
"billing_cycle": "yearly",
|
||||
"name": "高级会员年卡",
|
||||
"description": "高级会员年度订阅,享受全部功能",
|
||||
"price_cents": 39900,
|
||||
"monthly_price_cents": 3325,
|
||||
"duration_days": 365,
|
||||
"features": {
|
||||
"4K 超清分辨率": True,
|
||||
"批量渲染": True,
|
||||
"优先处理队列": True,
|
||||
"credits_per_month": 500,
|
||||
},
|
||||
}
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_membership_plans",
|
||||
return_value=[real_plan],
|
||||
):
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
assert len(plans) == 1
|
||||
p0 = plans[0]
|
||||
assert p0["plan_id"] == "premium_yearly"
|
||||
assert p0["price_cents"] == 39900
|
||||
assert p0["duration_days"] == 365
|
||||
assert p0["features"]["4K 超清分辨率"] is True
|
||||
def test_plans_endpoint_returns_three_tiers(self):
|
||||
import os # noqa: F401 (used by _import_plans_fn)
|
||||
|
||||
def test_plans_endpoint_empty_when_all_disabled(self):
|
||||
"""后台停用全部套餐时,用户端返回空列表。"""
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
with patch(
|
||||
"packages.application.catalog.admin_catalog.get_membership_plans",
|
||||
return_value=[],
|
||||
):
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
assert resp["plans"] == []
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
plan_ids = {p["plan_id"] for p in plans}
|
||||
assert plan_ids == {"monthly", "quarterly", "yearly"}
|
||||
for p in plans:
|
||||
assert p["price_cents"] > 0
|
||||
assert p["duration_days"] in (30, 90, 365)
|
||||
assert 0 < p["points_discount"] <= 1.0
|
||||
assert "max_resolution" in p["features"]
|
||||
|
||||
def test_longer_plans_cheaper_per_month(self):
|
||||
import os # noqa: F401
|
||||
|
||||
list_membership_plans = self._import_plans_fn()
|
||||
resp = list_membership_plans(current_user=_make_cu())
|
||||
plans = resp["plans"]
|
||||
monthly = next(p for p in plans if p["plan_id"] == "monthly")
|
||||
quarterly = next(p for p in plans if p["plan_id"] == "quarterly")
|
||||
yearly = next(p for p in plans if p["plan_id"] == "yearly")
|
||||
assert monthly["monthly_price_cents"] == 1990
|
||||
assert quarterly["monthly_price_cents"] < monthly["monthly_price_cents"]
|
||||
assert yearly["monthly_price_cents"] < quarterly["monthly_price_cents"]
|
||||
|
||||
|
||||
# ── P1-7: multiplier consistency ──────────────────────────────────────
|
||||
|
||||
|
||||
class TestMultiplierConsistency:
|
||||
def test_free_user_voice_clone_synth_1min_costs_2(self):
|
||||
"""voice_clone_synth base=1,免费用户 ceil(1*1.15)=2。"""
|
||||
def test_free_user_ai_title_costs_2(self):
|
||||
"""ai_title base=1,免费用户 ceil(1*1.15)=2。"""
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
assert calculate_points_cost("voice_clone_synth", is_member=False, duration_minutes=1) == 2
|
||||
|
||||
def test_retired_scenes_return_zero(self):
|
||||
"""已下线场景(ai_voice/ai_title/ai_cover/ai_rewrite 等)calculate_points_cost 统一返回 0。"""
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
for scene in ["ai_voice", "ai_title", "ai_cover", "ai_rewrite", "ai_video", "ai_digital_human"]:
|
||||
assert calculate_points_cost(scene, is_member=False, quantity=1) == 0
|
||||
assert calculate_points_cost("ai_title", is_member=False, quantity=1) == 2
|
||||
|
||||
def test_check_matches_direct_calculation(self):
|
||||
"""check 端点 required_points 与 calculate_points_cost 结果一致。"""
|
||||
@@ -239,14 +216,15 @@ class TestMultiplierConsistency:
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
svc = MagicMock()
|
||||
svc.check_daily_free_clip.return_value = False
|
||||
svc.get_or_create_account.return_value = {"balance": 999}
|
||||
db = MagicMock()
|
||||
cu = _make_cu()
|
||||
|
||||
with patch("app.api.routes.points._credits_enabled", return_value=True):
|
||||
for scene in ["voice_clone_synth", "voice_clone_train", "ai_voice", "ai_video", "ai_title"]:
|
||||
body = PointsCheckRequest(scene_key=scene, quantity=1, duration_minutes=1)
|
||||
for scene in ["ai_voice", "ai_title", "ai_cover", "ai_rewrite"]:
|
||||
body = PointsCheckRequest(scene_key=scene, quantity=1)
|
||||
with patch("app.api.routes.points._get_service", return_value=svc):
|
||||
resp = check_points(body=body, current_user=cu, db=db)
|
||||
expected = calculate_points_cost(scene, is_member=False, quantity=1, duration_minutes=1)
|
||||
expected = calculate_points_cost(scene, is_member=False, quantity=1)
|
||||
assert resp.required_points == expected, f"{scene}: got {resp.required_points}, expected {expected}"
|
||||
|
||||
+62
-565
@@ -1,4 +1,4 @@
|
||||
"""积分消耗规则单元测试 (#1895) — v1.6.2: 仅保留 voice_clone 相关"""
|
||||
"""积分消耗规则单元测试 (#1895)"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
@@ -7,6 +7,7 @@ import math
|
||||
import pytest
|
||||
|
||||
from packages.domain.points_rules import (
|
||||
DAILY_FREE_CLIP_LIMIT,
|
||||
FREE_USER_MULTIPLIER,
|
||||
MEMBER_DISCOUNT,
|
||||
MEMBERSHIP_PRICES,
|
||||
@@ -19,22 +20,8 @@ from packages.domain.points_rules import (
|
||||
class TestPointsScenesConfig:
|
||||
"""场景配置完整性"""
|
||||
|
||||
def test_registered_scenes_include_voice_clone_and_viral_video(self):
|
||||
"""场景配置:包含声音克隆(训练/合成)+ 爆款视频(动态定价)。"""
|
||||
assert {"voice_clone_train", "voice_clone_synth", "viral_video"}.issubset(set(POINTS_SCENES.keys()))
|
||||
|
||||
def test_viral_video_scene_is_dynamic_with_zero_base(self):
|
||||
"""viral_video 必须注册但 base_points=0 且 dynamic=True,不使用 @points_gate。"""
|
||||
vv = POINTS_SCENES["viral_video"]
|
||||
assert vv["base_points"] == 0
|
||||
assert vv["dynamic"] is True
|
||||
assert vv["unit"] == "次"
|
||||
assert vv["name"] == "爆款视频"
|
||||
|
||||
def test_voice_clone_scenes_defined(self):
|
||||
# 保留声音克隆两个场景
|
||||
assert "voice_clone_train" in POINTS_SCENES
|
||||
assert "voice_clone_synth" in POINTS_SCENES
|
||||
def test_all_nine_scenes_defined(self):
|
||||
assert len(POINTS_SCENES) == 9
|
||||
|
||||
def test_required_keys_present(self):
|
||||
for key, scene in POINTS_SCENES.items():
|
||||
@@ -45,14 +32,8 @@ class TestPointsScenesConfig:
|
||||
def test_voice_clone_train_is_free(self):
|
||||
assert POINTS_SCENES["voice_clone_train"]["base_points"] == 0
|
||||
|
||||
def test_voice_clone_synth_is_per_minute(self):
|
||||
assert POINTS_SCENES["voice_clone_synth"]["base_points"] == 1
|
||||
assert POINTS_SCENES["voice_clone_synth"]["unit"] == "分钟"
|
||||
|
||||
def test_calculate_points_cost_returns_zero_for_dynamic_viral_video(self):
|
||||
"""calculate_points_cost 对动态场景 viral_video 必须返回 0(由业务侧手动计算)。"""
|
||||
assert calculate_points_cost("viral_video", is_member=False) == 0.0
|
||||
assert calculate_points_cost("viral_video", is_member=True, member_type="monthly") == 0.0
|
||||
def test_ai_video_has_extra_per_30s(self):
|
||||
assert POINTS_SCENES["ai_video"]["extra_per_30s"] == 1
|
||||
|
||||
|
||||
class TestPointsPackages:
|
||||
@@ -71,23 +52,43 @@ class TestMembershipPrices:
|
||||
assert MEMBERSHIP_PRICES["yearly"]["price_cents"] == 15900
|
||||
|
||||
|
||||
class TestDailyFreeLimit:
|
||||
def test_limit_is_2(self):
|
||||
assert DAILY_FREE_CLIP_LIMIT == 2
|
||||
|
||||
|
||||
class TestCalculatePointsCost:
|
||||
"""核心计费逻辑"""
|
||||
|
||||
# ── 声音克隆合成(按时长计费) ──
|
||||
# ── 按次计费 ──
|
||||
|
||||
def test_voice_clone_synth_base(self):
|
||||
cost = calculate_points_cost("voice_clone_synth", is_member=False, duration_minutes=3)
|
||||
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
|
||||
|
||||
def test_voice_clone_synth_rounds_up(self):
|
||||
cost = calculate_points_cost("voice_clone_synth", is_member=False, duration_minutes=2.3)
|
||||
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
|
||||
|
||||
def test_voice_clone_synth_minimum_1_minute(self):
|
||||
cost = calculate_points_cost("voice_clone_synth", is_member=False, duration_minutes=0.1)
|
||||
def test_per_time_base_cost(self):
|
||||
# ai_rewrite: 1积分/次,免费用户 ceil(1 * 1.15) = 2
|
||||
cost = calculate_points_cost("ai_rewrite", is_member=False, quantity=1)
|
||||
assert cost == math.ceil(1 * FREE_USER_MULTIPLIER)
|
||||
|
||||
def test_per_time_multiple(self):
|
||||
# ai_cover: 1积分/张,3张 → base=3, free: ceil(3*1.15)=4
|
||||
cost = calculate_points_cost("ai_cover", is_member=False, quantity=3)
|
||||
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
|
||||
|
||||
# ── 按时长计费 ──
|
||||
|
||||
def test_per_minute_base(self):
|
||||
# ai_voice: 1积分/分钟,3分钟 → base=3, free: ceil(3*1.15)=4
|
||||
cost = calculate_points_cost("ai_voice", is_member=False, duration_minutes=3)
|
||||
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
|
||||
|
||||
def test_per_minute_rounds_up(self):
|
||||
# 2.3分钟 → ceil(2.3)=3分钟 → base=3
|
||||
cost = calculate_points_cost("ai_voice", is_member=False, duration_minutes=2.3)
|
||||
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
|
||||
|
||||
def test_digital_human_expensive(self):
|
||||
# ai_digital_human: 15积分/分钟,1分钟 → base=15, free: ceil(15*1.15)=18
|
||||
cost = calculate_points_cost("ai_digital_human", is_member=False, duration_minutes=1)
|
||||
assert cost == 18
|
||||
|
||||
# ── 免费场景 ──
|
||||
|
||||
def test_voice_clone_train_free(self):
|
||||
@@ -98,546 +99,42 @@ class TestCalculatePointsCost:
|
||||
cost = calculate_points_cost("voice_clone_train", is_member=True)
|
||||
assert cost == 0
|
||||
|
||||
# ── 混剪额外逻辑 ──
|
||||
|
||||
def test_ai_video_short_no_extra(self):
|
||||
# 20s (0.33min) ≤ 30s,不额外加积分,base=3, free: ceil(3*1.15)=4
|
||||
cost = calculate_points_cost("ai_video", is_member=False, quantity=1, duration_minutes=0.33)
|
||||
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
|
||||
|
||||
def test_ai_video_long_extra_charge(self):
|
||||
# 80s → base=3 + extra ceil((80-30)/30)=2 → total_base=5, free: ceil(5*1.15)=6
|
||||
cost = calculate_points_cost("ai_video", is_member=False, quantity=1, duration_minutes=80 / 60)
|
||||
assert cost == math.ceil(5 * FREE_USER_MULTIPLIER)
|
||||
|
||||
# ── 会员折扣 ──
|
||||
|
||||
def test_monthly_member_discount(self):
|
||||
cost = calculate_points_cost("voice_clone_synth", is_member=True, duration_minutes=1, member_type="monthly")
|
||||
assert cost == max(1, math.floor(1 * MEMBER_DISCOUNT["monthly"]))
|
||||
# ai_voice 1分钟 base=1, 月卡0.9 → floor(1*0.9)=1 → max(1,1)=1
|
||||
cost = calculate_points_cost("ai_voice", is_member=True, duration_minutes=1, member_type="monthly")
|
||||
assert cost == max(1, math.floor(1 * 0.9))
|
||||
|
||||
def test_yearly_member_deep_discount(self):
|
||||
# ai_digital_human 2分钟 base=30, 年卡0.8 → floor(30*0.8)=24
|
||||
cost = calculate_points_cost(
|
||||
"voice_clone_synth",
|
||||
"ai_digital_human",
|
||||
is_member=True,
|
||||
duration_minutes=2,
|
||||
member_type="yearly",
|
||||
)
|
||||
assert cost == max(1, math.floor(2 * MEMBER_DISCOUNT["yearly"]))
|
||||
assert cost == max(1, math.floor(30 * 0.8))
|
||||
|
||||
def test_member_without_type_no_discount(self):
|
||||
cost = calculate_points_cost("voice_clone_synth", is_member=True, duration_minutes=1)
|
||||
assert cost == 1
|
||||
# is_member=True 但没传 member_type → 不按会员折扣
|
||||
cost = calculate_points_cost("ai_voice", is_member=True, duration_minutes=1)
|
||||
assert cost == 1 # base=1, no discount applied
|
||||
|
||||
# ── 已下线/未知场景(向后兼容:返回 0) ──
|
||||
# ── 异常 ──
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"scene",
|
||||
[
|
||||
"ai_voice",
|
||||
"ai_video",
|
||||
"ai_digital_human",
|
||||
"ai_rewrite",
|
||||
"ai_cover",
|
||||
"ai_title",
|
||||
"douyin_extract",
|
||||
"nonexistent",
|
||||
],
|
||||
)
|
||||
def test_retired_scenes_return_zero(self, scene):
|
||||
assert calculate_points_cost(scene, is_member=False) == 0
|
||||
assert calculate_points_cost(scene, is_member=True, duration_minutes=10) == 0
|
||||
|
||||
|
||||
# ============ 爆款视频动态定价 (#2151) ============
|
||||
|
||||
|
||||
class TestResolveVideoDimensions:
|
||||
"""resolve_video_dimensions(): 分辨率别名、比例、默认兜底。"""
|
||||
|
||||
def test_1080p_16_9(self):
|
||||
"""1080p + 16:9 → w=1920, h=1080。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("1080p", "16:9")
|
||||
assert (w, h) == (1920, 1080)
|
||||
|
||||
def test_480p_16_9(self):
|
||||
"""480p + 16:9 → 854×480(ceil(480*16/9)=854,偶对齐)。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("480p", "16:9")
|
||||
assert (w, h) == (854, 480)
|
||||
|
||||
def test_720p_1_1(self):
|
||||
"""1:1 正方形 → w == h。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("720p", "1:1")
|
||||
assert (w, h) == (720, 720)
|
||||
|
||||
def test_1080p_1_1(self):
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("1080p", "1:1")
|
||||
assert (w, h) == (1080, 1080)
|
||||
|
||||
def test_resolution_aliases(self):
|
||||
"""中文/英文别名应正确映射到对应高度。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
cases = [
|
||||
("普清", 480),
|
||||
("sd", 480),
|
||||
("low", 480),
|
||||
("default", 480),
|
||||
("高清", 720),
|
||||
("medium", 720),
|
||||
("hd", 720),
|
||||
("超清", 1080),
|
||||
("fhd", 1080),
|
||||
("ultra", 1080),
|
||||
("全能", 1080),
|
||||
("high", 1080),
|
||||
]
|
||||
for alias, expected_h in cases:
|
||||
_, h = resolve_video_dimensions(alias, "1:1")
|
||||
assert h == expected_h, f"{alias} -> h={h}, expected {expected_h}"
|
||||
|
||||
def test_unknown_resolution_falls_back_to_720p(self):
|
||||
"""未知分辨率字符串兜底到 720p。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("garbage-xxx", "1:1")
|
||||
assert h == 720
|
||||
assert w == 720
|
||||
|
||||
def test_4k_16_9(self):
|
||||
"""#2159 4k 横屏:短边=height=2160,width=3840。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("4k", "16:9")
|
||||
assert (w, h) == (3840, 2160)
|
||||
|
||||
def test_2160p_alias(self):
|
||||
"""2160p 别名→4k。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("2160p", "9:16")
|
||||
assert (w, h) == (2160, 3840)
|
||||
|
||||
def test_empty_resolution_defaults_to_720p_9_16(self):
|
||||
"""空 resolution + 空 ratio → 默认 720p + 9:16 竖屏 (720×1280)。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("", "")
|
||||
assert (w, h) == (720, 1280)
|
||||
|
||||
def test_none_resolution_default_ratio(self):
|
||||
"""None resolution + None ratio → 720p + 9:16 竖屏默认。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions(None, None)
|
||||
assert (w, h) == (720, 1280)
|
||||
|
||||
def test_720p_9_16_portrait(self):
|
||||
"""720p + 9:16 竖屏 → 短边是 width=720,height=1280(v10实测)。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("720p", "9:16")
|
||||
assert (w, h) == (720, 1280)
|
||||
|
||||
def test_1080p_9_16_portrait(self):
|
||||
"""1080p + 9:16 竖屏 → 1080×1920。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("1080p", "9:16")
|
||||
assert (w, h) == (1080, 1920)
|
||||
|
||||
def test_480p_9_16_portrait(self):
|
||||
"""480p + 9:16 竖屏 → 480×854。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions("480p", "9:16")
|
||||
assert (w, h) == (480, 854)
|
||||
|
||||
def test_all_dimensions_even(self):
|
||||
"""所有返回尺寸都应是偶数(视频编码要求)。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
for res in ("480p", "720p", "1080p", "普清", "高清", "超清"):
|
||||
for ratio in ("16:9", "9:16", "1:1"):
|
||||
w, h = resolve_video_dimensions(res, ratio)
|
||||
assert w % 2 == 0 and h % 2 == 0, f"{res}/{ratio} -> ({w},{h}) not even"
|
||||
|
||||
def test_whitespace_resolution_case_insensitive(self):
|
||||
"""前后空格 + 大写应被规范化处理。"""
|
||||
from packages.domain.points_rules import resolve_video_dimensions
|
||||
|
||||
w, h = resolve_video_dimensions(" 1080P ", " 16:9 ")
|
||||
assert (w, h) == (1920, 1080)
|
||||
|
||||
|
||||
class TestMatchModelPrefix:
|
||||
"""_match_model_prefix() 前缀匹配 + 兜底。"""
|
||||
|
||||
def test_seedance_2_5_exact(self):
|
||||
from packages.domain.points_rules import _match_model_prefix
|
||||
|
||||
assert _match_model_prefix("seedance-2.5") == "seedance-2.5"
|
||||
|
||||
def test_seedance_2_5_with_variant(self):
|
||||
"""带后缀版本号(如 seedance-2.5-pro)仍匹配 seedance-2.5。"""
|
||||
from packages.domain.points_rules import _match_model_prefix
|
||||
|
||||
assert _match_model_prefix("seedance-2.5-pro") == "seedance-2.5"
|
||||
|
||||
def test_seedance_2_0_exact(self):
|
||||
from packages.domain.points_rules import _match_model_prefix
|
||||
|
||||
assert _match_model_prefix("seedance-2.0") == "seedance-2.0"
|
||||
|
||||
def test_seedance_2_0_with_variant(self):
|
||||
from packages.domain.points_rules import _match_model_prefix
|
||||
|
||||
assert _match_model_prefix("seedance-2.0-lite") == "seedance-2.0"
|
||||
|
||||
def test_unknown_model_falls_back_to_2_5(self):
|
||||
"""未知模型前缀兜底 seedance-2.5。"""
|
||||
from packages.domain.points_rules import _match_model_prefix
|
||||
|
||||
assert _match_model_prefix("kling-v1") == "seedance-2.5"
|
||||
assert _match_model_prefix("") == "seedance-2.5"
|
||||
assert _match_model_prefix(None) == "seedance-2.5"
|
||||
|
||||
def test_case_insensitive(self):
|
||||
from packages.domain.points_rules import _match_model_prefix
|
||||
|
||||
assert _match_model_prefix("SEEDANCE-2.0") == "seedance-2.0"
|
||||
|
||||
|
||||
class TestInferResolutionKey:
|
||||
"""_infer_resolution_key(w, h): 按短边 1000+/650-999/<650 三档。"""
|
||||
|
||||
def test_short_side_ge_1000_is_1080p(self):
|
||||
from packages.domain.points_rules import _infer_resolution_key
|
||||
|
||||
assert _infer_resolution_key(1920, 1080) == "1080p" # 横屏
|
||||
assert _infer_resolution_key(1080, 1920) == "1080p" # 竖屏
|
||||
assert _infer_resolution_key(1080, 1080) == "1080p" # 方屏
|
||||
|
||||
def test_short_side_650_to_999_is_720p(self):
|
||||
from packages.domain.points_rules import _infer_resolution_key
|
||||
|
||||
assert _infer_resolution_key(1280, 720) == "720p"
|
||||
assert _infer_resolution_key(720, 1280) == "720p"
|
||||
assert _infer_resolution_key(720, 720) == "720p"
|
||||
|
||||
def test_short_side_lt_650_is_480p(self):
|
||||
from packages.domain.points_rules import _infer_resolution_key
|
||||
|
||||
assert _infer_resolution_key(854, 480) == "480p"
|
||||
assert _infer_resolution_key(480, 854) == "480p"
|
||||
assert _infer_resolution_key(480, 480) == "480p"
|
||||
# 极小值兜底
|
||||
assert _infer_resolution_key(1, 1) == "480p"
|
||||
|
||||
def test_portrait_1280_height_is_720p_short_side(self):
|
||||
"""竖屏 720×1280 短边=720,应识别为 720p 而非 1080p(老bug回归)。"""
|
||||
from packages.domain.points_rules import _infer_resolution_key
|
||||
|
||||
assert _infer_resolution_key(720, 1280) == "720p"
|
||||
|
||||
|
||||
class TestCalculateViralVideoCredits:
|
||||
"""calculate_viral_video_credits():爆款视频动态定价核心函数。"""
|
||||
|
||||
def test_default_args_returns_float(self):
|
||||
"""默认参数返回 float。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
credits = calculate_viral_video_credits(15, 1280, 720)
|
||||
assert isinstance(credits, float)
|
||||
|
||||
def test_return_is_rounded_to_two_decimals(self):
|
||||
"""round(..., 2) 后值本身就是两位小数(再 round 不变化)。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
for dur, w, h in [(15, 1280, 720), (5, 854, 480), (30, 1920, 1080), (10, 720, 720)]:
|
||||
credits = calculate_viral_video_credits(dur, w, h)
|
||||
assert round(credits, 2) == credits
|
||||
|
||||
def test_has_video_input_uses_lower_price(self):
|
||||
"""has_video_input=True 时使用参考视频价格(有视频输入便宜)。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
no_input = calculate_viral_video_credits(15, 1280, 720, model="seedance-2.5", has_video_input=False)
|
||||
with_input = calculate_viral_video_credits(15, 1280, 720, model="seedance-2.5", has_video_input=True)
|
||||
assert with_input < no_input
|
||||
|
||||
def test_unknown_model_falls_back_to_seedance_2_5(self):
|
||||
"""未知 model 前缀兜底到 seedance-2.5 价格,与默认等价。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
unknown = calculate_viral_video_credits(15, 1280, 720, model="unknown-model")
|
||||
default = calculate_viral_video_credits(15, 1280, 720, model="seedance-2.5")
|
||||
assert unknown == default
|
||||
|
||||
def test_actual_tokens_overrides_calculation(self):
|
||||
"""传入 actual_tokens>0 时用它替代公式计算的 tokens。"""
|
||||
from packages.domain.points_rules import (
|
||||
VIRAL_VIDEO_FIXED_COST,
|
||||
VIRAL_VIDEO_MODEL_PRICES,
|
||||
VIRAL_VIDEO_PROFIT_MULTIPLIER,
|
||||
calculate_viral_video_credits,
|
||||
)
|
||||
|
||||
price = VIRAL_VIDEO_MODEL_PRICES[("seedance-2.5", "720p", False)]
|
||||
actual_tokens = 2_000_000
|
||||
expected = round(
|
||||
(actual_tokens / 1_000_000 * price + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER, 2
|
||||
)
|
||||
credits = calculate_viral_video_credits(15, 1280, 720, actual_tokens=actual_tokens)
|
||||
assert credits == expected
|
||||
|
||||
def test_zero_duration_width_height_defensive_max1(self):
|
||||
"""duration/width/height 为 0/None 时 max(1,...) 防御,结果>0。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c_zero = calculate_viral_video_credits(0, 0, 0)
|
||||
assert c_zero > 0
|
||||
c_none = calculate_viral_video_credits(None, None, None)
|
||||
assert c_none > 0
|
||||
c_one = calculate_viral_video_credits(1, 1, 1)
|
||||
assert c_none == c_one
|
||||
|
||||
def test_non_default_fps_affects_tokens(self):
|
||||
"""fps 非默认值(30) 应比默认(24) 积分高。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c24 = calculate_viral_video_credits(15, 1280, 720, fps=24)
|
||||
c30 = calculate_viral_video_credits(15, 1280, 720, fps=30)
|
||||
assert c30 > c24
|
||||
|
||||
def test_seedance_2_0_priced_lower_than_2_5_at_1080p(self):
|
||||
"""seedance-2.0 在 1080p 无视频输入时定价低于 seedance-2.5。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c20 = calculate_viral_video_credits(15, 1920, 1080, model="seedance-2.0", has_video_input=False)
|
||||
c25 = calculate_viral_video_credits(15, 1920, 1080, model="seedance-2.5", has_video_input=False)
|
||||
assert c20 < c25
|
||||
|
||||
def test_formula_includes_fixed_cost_and_multiplier(self):
|
||||
"""手算公式结果应与函数返回一致(固定成本 + 利润系数)。"""
|
||||
from packages.domain.points_rules import (
|
||||
VIRAL_VIDEO_FIXED_COST,
|
||||
VIRAL_VIDEO_FPS,
|
||||
VIRAL_VIDEO_MODEL_PRICES,
|
||||
VIRAL_VIDEO_PROFIT_MULTIPLIER,
|
||||
calculate_viral_video_credits,
|
||||
)
|
||||
|
||||
dur, w, h = 10, 1280, 720
|
||||
price = VIRAL_VIDEO_MODEL_PRICES[("seedance-2.5", "720p", False)]
|
||||
tokens = dur * w * h * VIRAL_VIDEO_FPS / 1024.0
|
||||
expected = round((tokens / 1_000_000 * price + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER, 2)
|
||||
assert calculate_viral_video_credits(dur, w, h) == expected
|
||||
|
||||
def test_seedance_2_0_with_video_input_falls_back_to_seedance_2_5_price(self):
|
||||
"""seedance-2.0 + has_video_input=True 组合不在价格表,走 line 111 fallback 到 seedance-2.5 的 720p False 价格。"""
|
||||
from packages.domain.points_rules import (
|
||||
VIRAL_VIDEO_FIXED_COST,
|
||||
VIRAL_VIDEO_FPS,
|
||||
VIRAL_VIDEO_MODEL_PRICES,
|
||||
VIRAL_VIDEO_PROFIT_MULTIPLIER,
|
||||
calculate_viral_video_credits,
|
||||
)
|
||||
|
||||
dur, w, h = 10, 1280, 720
|
||||
# 兜底价格 = seedance-2.5/720p/False = 70.0
|
||||
price = VIRAL_VIDEO_MODEL_PRICES[("seedance-2.5", "720p", False)]
|
||||
assert price == 70.0
|
||||
tokens = dur * w * h * VIRAL_VIDEO_FPS / 1024.0
|
||||
expected = round((tokens / 1_000_000 * price + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER, 2)
|
||||
credits = calculate_viral_video_credits(dur, w, h, model="seedance-2.0", has_video_input=True)
|
||||
assert credits == expected
|
||||
|
||||
def test_fps_zero_or_none_falls_back_to_default(self):
|
||||
"""fps=0/None 时 int(fps or 24) 兜底到默认 24,结果与 fps=24 一致。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c_default = calculate_viral_video_credits(10, 1280, 720, fps=24)
|
||||
c_zero = calculate_viral_video_credits(10, 1280, 720, fps=0)
|
||||
c_none = calculate_viral_video_credits(10, 1280, 720, fps=None)
|
||||
assert c_zero == c_default
|
||||
assert c_none == c_default
|
||||
|
||||
# ──────── P0 计费回归:短边规则 + 价格精确断言 ────────
|
||||
|
||||
def test_15s_720p_portrait_is_29_68(self):
|
||||
"""P0 回归:15s/720p/9:16 竖屏 (720×1280) 必须 =29.68 积分。"""
|
||||
from packages.domain.points_rules import (
|
||||
calculate_viral_video_credits,
|
||||
resolve_video_dimensions,
|
||||
)
|
||||
|
||||
w, h = resolve_video_dimensions("720p", "9:16")
|
||||
assert (w, h) == (720, 1280)
|
||||
assert calculate_viral_video_credits(15, w, h) == 29.68
|
||||
|
||||
def test_30s_1080p_portrait_is_146_14(self):
|
||||
"""P0 回归:30s/1080p/9:16 竖屏 (1080×1920) =146.14 积分。"""
|
||||
from packages.domain.points_rules import (
|
||||
calculate_viral_video_credits,
|
||||
resolve_video_dimensions,
|
||||
)
|
||||
|
||||
w, h = resolve_video_dimensions("1080p", "9:16")
|
||||
assert (w, h) == (1080, 1920)
|
||||
assert calculate_viral_video_credits(30, w, h) == 146.14
|
||||
|
||||
def test_portrait_landscape_same_pixels_same_price(self):
|
||||
"""相同像素数(横竖屏旋转)积分一致。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
assert calculate_viral_video_credits(15, 1280, 720) == calculate_viral_video_credits(15, 720, 1280)
|
||||
|
||||
|
||||
class TestViralVideoCreditsWithBreakdown:
|
||||
"""calculate_viral_video_credits_with_breakdown:返回 (credits, breakdown_dict)。"""
|
||||
|
||||
def test_returns_credits_matching_plain_version(self):
|
||||
"""新函数返回的 credits 必须与 calculate_viral_video_credits 完全一致,且 breakdown 字段齐全。"""
|
||||
from packages.domain.points_rules import (
|
||||
calculate_viral_video_credits,
|
||||
calculate_viral_video_credits_with_breakdown,
|
||||
)
|
||||
|
||||
for dur, w, h, model, hvi in [
|
||||
(15, 1280, 720, "seedance-2.5", False),
|
||||
(10, 720, 1280, "seedance-2.0", False),
|
||||
(30, 1920, 1080, "seedance-2.5", False),
|
||||
(5, 480, 480, "", False),
|
||||
]:
|
||||
c1 = calculate_viral_video_credits(dur, w, h, model=model, has_video_input=hvi)
|
||||
c2, bd = calculate_viral_video_credits_with_breakdown(dur, w, h, model=model, has_video_input=hvi)
|
||||
assert c1 == c2
|
||||
assert isinstance(bd, dict)
|
||||
for key in (
|
||||
"tokens",
|
||||
"video_cost",
|
||||
"fixed_cost",
|
||||
"profit_multiplier",
|
||||
"model_price",
|
||||
"width",
|
||||
"height",
|
||||
"fps",
|
||||
):
|
||||
assert key in bd, f"breakdown missing key: {key}"
|
||||
assert bd["fixed_cost"] == 0.15
|
||||
assert bd["profit_multiplier"] == 1.3
|
||||
assert bd["width"] == w
|
||||
assert bd["height"] == h
|
||||
assert bd["fps"] == 24
|
||||
assert bd["tokens"] > 0
|
||||
assert bd["model_price"] > 0
|
||||
expected = round((bd["video_cost"] + bd["fixed_cost"]) * bd["profit_multiplier"], 2)
|
||||
assert expected == c2
|
||||
|
||||
def test_actual_tokens_overrides_computed(self):
|
||||
"""actual_tokens 传入时应覆盖按公式计算的 tokens。"""
|
||||
from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown
|
||||
|
||||
c, bd = calculate_viral_video_credits_with_breakdown(
|
||||
15,
|
||||
1280,
|
||||
720,
|
||||
actual_tokens=1_000_000,
|
||||
)
|
||||
assert bd["tokens"] == 1_000_000.0
|
||||
# video_cost = 1M/1M * 70 = 70; total = (70+0.15)*1.3 = 91.195 → 91.20
|
||||
assert c == 91.20
|
||||
|
||||
|
||||
# ============ #2159 多模型定价单测 ============
|
||||
|
||||
|
||||
class TestMultiModelCredits:
|
||||
"""#2159 多模型积分估算正确性(含 token/second 两种计费模式)。"""
|
||||
|
||||
def test_seedance_2_5_15s_720p_9x16(self):
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
# 15s/720p/9:16 → 720×1280
|
||||
# tokens = 15*720*1280*24/1024 = 324000
|
||||
# video_cost = 324000/1M*70 = 22.68
|
||||
# total = (22.68+0.15)*1.3 = 29.679 ≈ 29.68
|
||||
c = calculate_viral_video_credits(15, 720, 1280, model="seedance-2.5")
|
||||
assert c == 29.68, f"got {c}"
|
||||
|
||||
def test_seedance_2_0_30s_1080p_9x16(self):
|
||||
# 30s/1080p/9:16 → 1080×1920
|
||||
# tokens = 30*1080*1920*24/1024 = 1,458,000
|
||||
# video_cost = 1.458M/1M*51 = 74.358
|
||||
# total = (74.358+0.15)*1.3 = 96.86
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(30, 1080, 1920, model="seedance-2.0")
|
||||
assert c == 96.86, f"got {c}"
|
||||
|
||||
def test_seedance_2_0_fast_15s_720p_9x16(self):
|
||||
# 15s/720p/9:16 tokens=324000, price=28
|
||||
# video_cost = 0.324*28 = 9.072
|
||||
# total = (9.072+0.15)*1.3 = 11.99
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(15, 720, 1280, model="seedance-2.0-fast")
|
||||
assert c == 11.99, f"got {c}"
|
||||
|
||||
def test_seedance_2_0_mini_15s_720p_9x16(self):
|
||||
# price=9.2, tokens=324000
|
||||
# video_cost = 0.324*9.2 = 2.9808
|
||||
# total = (2.9808+0.15)*1.3 = 4.07
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(15, 720, 1280, model="seedance-2.0-mini")
|
||||
assert c == 4.07, f"got {c}"
|
||||
|
||||
def test_wan_3_0_per_second_billing(self):
|
||||
# per_second: 10s/720p price=0.6元/秒
|
||||
# video_cost = 10*0.6 = 6.0
|
||||
# total = (6.0+0.15)*1.3 = 7.995 ≈ 8.00
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(10, 720, 1280, model="wan-3.0")
|
||||
assert c == 8.0, f"got {c}"
|
||||
|
||||
def test_seedance_2_0_4k_16x9(self):
|
||||
# 5s/4k/16:9 → 3840×2160, price=80
|
||||
# tokens = 5*3840*2160*24/1024 = 972000
|
||||
# video_cost = 0.972*80 = 77.76
|
||||
# total = (77.76+0.15)*1.3 = 101.28
|
||||
from packages.domain.points_rules import calculate_viral_video_credits
|
||||
|
||||
c = calculate_viral_video_credits(5, 3840, 2160, model="seedance-2.0")
|
||||
assert c == 101.28, f"got {c}"
|
||||
|
||||
def test_model_config_has_all_6_models(self):
|
||||
from packages.domain.points_rules import VIRAL_VIDEO_MODEL_CONFIG
|
||||
|
||||
expected = {"seedance-2.5", "seedance-2.0", "seedance-2.0-fast", "seedance-2.0-mini", "wan-3.0"}
|
||||
assert expected.issubset(set(VIRAL_VIDEO_MODEL_CONFIG.keys()))
|
||||
|
||||
def test_list_models_hides_wan_when_dashscope_unavailable(self):
|
||||
from packages.domain.points_rules import list_viral_video_models
|
||||
|
||||
all_models = list_viral_video_models(include_placeholder=False, dashscope_available=False)
|
||||
keys = {m["key"] for m in all_models}
|
||||
assert "wan-3.0" not in keys
|
||||
assert "seedance-2.5" in keys
|
||||
# is_default
|
||||
defaults = [m for m in all_models if m["is_default"]]
|
||||
assert len(defaults) == 1
|
||||
assert defaults[0]["key"] == "seedance-2.5"
|
||||
|
||||
def test_list_models_includes_wan_when_dashscope_available(self):
|
||||
from packages.domain.points_rules import list_viral_video_models
|
||||
|
||||
models = list_viral_video_models(include_placeholder=False, dashscope_available=True)
|
||||
keys = {m["key"] for m in models}
|
||||
assert "wan-3.0" in keys
|
||||
|
||||
def test_infer_4k(self):
|
||||
from packages.domain.points_rules import _infer_resolution_key
|
||||
|
||||
assert _infer_resolution_key(3840, 2160) == "4k"
|
||||
assert _infer_resolution_key(2160, 3840) == "4k"
|
||||
assert _infer_resolution_key(1920, 1080) == "1080p"
|
||||
def test_unknown_scene_raises(self):
|
||||
with pytest.raises(ValueError, match="Unknown points scene"):
|
||||
calculate_points_cost("nonexistent_scene", is_member=False)
|
||||
|
||||
@@ -72,7 +72,7 @@ class TestCheckBalance:
|
||||
|
||||
class TestDeductPoints:
|
||||
def test_deduct_fails_insufficient_balance(self, service, db_session, user_id):
|
||||
result = service.deduct_points(user_id, 100, "voice_clone_synth", db_session)
|
||||
result = service.deduct_points(user_id, 100, "ai_voice", db_session)
|
||||
assert result["success"] is False
|
||||
assert result["transaction_id"] is None
|
||||
|
||||
@@ -80,13 +80,13 @@ class TestDeductPoints:
|
||||
# 先充值
|
||||
service.add_points(user_id, 50, "recharge", db_session)
|
||||
# 再扣减
|
||||
result = service.deduct_points(user_id, 20, "voice_clone_synth", db_session)
|
||||
result = service.deduct_points(user_id, 20, "ai_voice", db_session)
|
||||
assert result["success"] is True
|
||||
assert result["balance"] == 30
|
||||
|
||||
def test_deduct_creates_transaction(self, service, db_session, user_id):
|
||||
service.add_points(user_id, 100, "recharge", db_session)
|
||||
result = service.deduct_points(user_id, 30, "voice_clone_synth", db_session)
|
||||
result = service.deduct_points(user_id, 30, "ai_voice", db_session)
|
||||
assert result["success"] is True
|
||||
|
||||
txns = service.get_transactions(user_id, db_session)
|
||||
@@ -111,14 +111,14 @@ class TestAddPoints:
|
||||
class TestRefundPoints:
|
||||
def test_refund_adds_back(self, service, db_session, user_id):
|
||||
service.add_points(user_id, 100, "recharge", db_session)
|
||||
service.deduct_points(user_id, 20, "voice_clone_synth", db_session)
|
||||
result = service.refund_points(user_id, 20, "voice_clone_synth", db_session)
|
||||
service.deduct_points(user_id, 20, "ai_voice", db_session)
|
||||
result = service.refund_points(user_id, 20, "ai_voice", db_session)
|
||||
assert result["success"] is True
|
||||
assert result["balance"] == 100
|
||||
|
||||
def test_refund_creates_refund_transaction(self, service, db_session, user_id):
|
||||
service.add_points(user_id, 100, "recharge", db_session)
|
||||
service.refund_points(user_id, 10, "voice_clone_synth", db_session)
|
||||
service.refund_points(user_id, 10, "ai_rewrite", db_session)
|
||||
|
||||
txns = service.get_transactions(user_id, db_session)
|
||||
refund_txns = [t for t in txns["items"] if t["type"] == "add" and "refund" in t["source"]]
|
||||
@@ -145,15 +145,21 @@ class TestGetTransactions:
|
||||
|
||||
|
||||
class TestGetDailyUsage:
|
||||
"""智能混剪已免费,get_daily_usage 返回 unlimited(-1)占位。"""
|
||||
|
||||
def test_returns_unlimited(self, service, db_session, user_id):
|
||||
result = service.get_daily_usage(user_id, db_session)
|
||||
def test_zero_usage(self, service, db_session, user_id):
|
||||
with patch("packages.domain.points_service._get_redis_client", return_value=None):
|
||||
result = service.get_daily_usage(user_id, db_session)
|
||||
assert result["free_clips_used"] == 0
|
||||
assert result["free_clips_limit"] == -1 # -1 表示 unlimited
|
||||
assert result["free_clips_remaining"] == -1
|
||||
assert result["free_clips_limit"] == 2
|
||||
assert result["free_clips_remaining"] == 2
|
||||
assert "reset_at" in result
|
||||
|
||||
def test_after_recording(self, service, db_session, user_id):
|
||||
with patch("packages.domain.points_service._get_redis_client", return_value=None):
|
||||
service.record_daily_free_clip(user_id, db_session)
|
||||
result = service.get_daily_usage(user_id, db_session)
|
||||
assert result["free_clips_used"] == 1
|
||||
assert result["free_clips_remaining"] == 1
|
||||
|
||||
|
||||
class TestCreateOrder:
|
||||
def test_points_order(self, service, db_session, user_id):
|
||||
@@ -179,172 +185,3 @@ class TestCreateOrder:
|
||||
def test_unknown_order_type_raises(self, service, db_session, user_id):
|
||||
with pytest.raises(ValueError, match="Unknown order type"):
|
||||
service.create_order(user_id, "insurance", "basic", db_session)
|
||||
|
||||
|
||||
# ============ 爆款视频(viral_video)动态定价方法 ============
|
||||
|
||||
|
||||
class TestDeductViralVideo:
|
||||
"""deduct_viral_video(): 预扣积分,委托给 deduct_points。"""
|
||||
|
||||
def test_delegates_to_deduct_points_with_correct_args(self, service, db_session, user_id):
|
||||
"""deduct_viral_video 应以 source='viral_video', ref_id=job_id 调用 deduct_points。"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
expected = {"success": True, "balance": 50.0, "transaction_id": "t1"}
|
||||
with patch.object(service, "deduct_points", return_value=expected) as mock_dp:
|
||||
result = service.deduct_viral_video(user_id, 10.5, "job-abc", db_session)
|
||||
|
||||
assert result == expected
|
||||
mock_dp.assert_called_once()
|
||||
kwargs = mock_dp.call_args.kwargs
|
||||
assert kwargs["user_id"] == user_id
|
||||
assert kwargs["amount"] == 10.5
|
||||
assert kwargs["source"] == "viral_video"
|
||||
assert kwargs["db"] is db_session
|
||||
assert kwargs["description"] == "爆款视频生成"
|
||||
assert kwargs["ref_id"] == "job-abc"
|
||||
|
||||
def test_none_credits_coerced_to_zero(self, service, db_session, user_id):
|
||||
"""credits=None 时应被 float(credits or 0) 转为 0,不抛异常。"""
|
||||
|
||||
with patch.object(
|
||||
service, "deduct_points", return_value={"success": True, "balance": 0, "transaction_id": "t"}
|
||||
) as mock_dp:
|
||||
service.deduct_viral_video(user_id, None, "job-nil", db_session)
|
||||
assert mock_dp.call_args.kwargs["amount"] == 0.0
|
||||
|
||||
|
||||
class TestSettleViralVideo:
|
||||
"""settle_viral_video(): 多退少补结算。"""
|
||||
|
||||
def test_no_action_when_diff_below_epsilon(self, service, db_session, user_id):
|
||||
"""|diff|<0.01 时返回 action=none,不调 refund/deduct。"""
|
||||
|
||||
with (
|
||||
patch.object(service, "refund_points") as mock_refund,
|
||||
patch.object(service, "deduct_points") as mock_deduct,
|
||||
):
|
||||
result = service.settle_viral_video(user_id, estimated=10.00, actual=10.001, txn_id="t1", db=db_session)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["action"] == "none"
|
||||
assert result["diff"] == 0.0
|
||||
mock_refund.assert_not_called()
|
||||
mock_deduct.assert_not_called()
|
||||
|
||||
def test_refund_when_actual_less_than_estimated(self, service, db_session, user_id):
|
||||
"""actual<estimated 时走 refund_points,返回 action=refund。"""
|
||||
|
||||
refund_res = {"success": True, "balance": 60.0, "transaction_id": "tr-1"}
|
||||
with patch.object(service, "refund_points", return_value=refund_res) as mock_refund:
|
||||
result = service.settle_viral_video(user_id, estimated=20.0, actual=15.0, txn_id="t2", db=db_session)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["action"] == "refund"
|
||||
assert result["amount"] == 5.0
|
||||
assert result["diff"] == -5.0
|
||||
mock_refund.assert_called_once()
|
||||
rk = mock_refund.call_args.kwargs
|
||||
assert rk["user_id"] == user_id
|
||||
assert rk["amount"] == 5.0
|
||||
assert rk["source"] == "viral_video"
|
||||
assert rk["ref_id"] == "t2"
|
||||
assert rk["description"] == "爆款视频结算退费"
|
||||
|
||||
def test_refund_exception_returns_failure(self, service, db_session, user_id):
|
||||
"""refund_points 抛异常时,应捕获并返回 success=False。"""
|
||||
|
||||
with patch.object(service, "refund_points", side_effect=RuntimeError("db down")):
|
||||
result = service.settle_viral_video(user_id, estimated=20.0, actual=10.0, txn_id="t3", db=db_session)
|
||||
|
||||
assert result["success"] is False
|
||||
assert result["action"] == "refund"
|
||||
|
||||
def test_deduct_when_actual_greater_than_estimated_success(self, service, db_session, user_id):
|
||||
"""actual>estimated 且补扣成功 → action=deduct, success=True。"""
|
||||
|
||||
deduct_res = {"success": True, "balance": 40.0, "transaction_id": "td-1"}
|
||||
with patch.object(service, "deduct_points", return_value=deduct_res) as mock_deduct:
|
||||
result = service.settle_viral_video(user_id, estimated=15.0, actual=20.0, txn_id="t4", db=db_session)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["action"] == "deduct"
|
||||
assert result["amount"] == 5.0
|
||||
assert result["diff"] == 5.0
|
||||
mock_deduct.assert_called_once()
|
||||
dk = mock_deduct.call_args.kwargs
|
||||
assert dk["amount"] == 5.0
|
||||
assert dk["source"] == "viral_video"
|
||||
assert dk["ref_id"] == "t4"
|
||||
|
||||
def test_deduct_insufficient_balance_returns_success_false_not_raise(self, service, db_session, user_id):
|
||||
"""actual>estimated 补扣时余额不足(success=False)应记录 warning 但不抛异常。"""
|
||||
import logging
|
||||
|
||||
deduct_res = {"success": False, "balance": 2.0, "transaction_id": None}
|
||||
with (
|
||||
patch.object(service, "deduct_points", return_value=deduct_res),
|
||||
patch("packages.domain.points_service.logger") as mock_logger,
|
||||
):
|
||||
result = service.settle_viral_video(user_id, estimated=15.0, actual=20.0, txn_id="t5", db=db_session)
|
||||
|
||||
# 即使补扣失败,函数也返回 action=deduct 但 success=False(不阻塞任务完成)
|
||||
assert result["success"] is False
|
||||
assert result["action"] == "deduct"
|
||||
assert result["amount"] == 5.0
|
||||
# 应打印 warning
|
||||
mock_logger.warning.assert_called_once()
|
||||
|
||||
def test_deduct_exception_returns_failure(self, service, db_session, user_id):
|
||||
"""deduct_points 抛异常时应捕获并返回 success=False。"""
|
||||
|
||||
with patch.object(service, "deduct_points", side_effect=RuntimeError("db boom")):
|
||||
result = service.settle_viral_video(user_id, estimated=10.0, actual=20.0, txn_id="t6", db=db_session)
|
||||
|
||||
assert result["success"] is False
|
||||
assert result["action"] == "deduct"
|
||||
assert result["diff"] == 10.0
|
||||
|
||||
|
||||
class TestRefundViralVideo:
|
||||
"""refund_viral_video(): 爆款视频失败全额退款。"""
|
||||
|
||||
def test_zero_amount_returns_none_action(self, service, db_session, user_id):
|
||||
"""amount<=0 直接返回 none action,不调 refund_points。"""
|
||||
|
||||
with patch.object(service, "refund_points") as mock_refund:
|
||||
r1 = service.refund_viral_video(user_id, 0, "t0", db_session)
|
||||
r2 = service.refund_viral_video(user_id, None, "t0", db_session)
|
||||
r3 = service.refund_viral_video(user_id, -1.5, "t0", db_session)
|
||||
|
||||
assert r1 == {"success": True, "action": "none", "amount": 0.0}
|
||||
assert r2 == {"success": True, "action": "none", "amount": 0.0}
|
||||
assert r3["action"] == "none"
|
||||
mock_refund.assert_not_called()
|
||||
|
||||
def test_success_path_delegates_to_refund_points(self, service, db_session, user_id):
|
||||
"""成功路径:透传 user_id/amount/ref_id=txn_id/source=viral_video。"""
|
||||
|
||||
expected = {"success": True, "balance": 80.0, "transaction_id": "rf-1"}
|
||||
with patch.object(service, "refund_points", return_value=expected) as mock_refund:
|
||||
result = service.refund_viral_video(user_id, 30.0, "txn-xyz", db_session)
|
||||
|
||||
assert result == expected
|
||||
mock_refund.assert_called_once()
|
||||
rk = mock_refund.call_args.kwargs
|
||||
assert rk["user_id"] == user_id
|
||||
assert rk["amount"] == 30.0
|
||||
assert rk["source"] == "viral_video"
|
||||
assert rk["ref_id"] == "txn-xyz"
|
||||
assert rk["description"] == "爆款视频失败退款"
|
||||
|
||||
def test_exception_returns_failure(self, service, db_session, user_id):
|
||||
"""refund_points 抛异常时返回 success=False/action=refund。"""
|
||||
|
||||
with patch.object(service, "refund_points", side_effect=RuntimeError("conn lost")):
|
||||
result = service.refund_viral_video(user_id, 25.0, "txn-err", db_session)
|
||||
|
||||
assert result["success"] is False
|
||||
assert result["action"] == "refund"
|
||||
assert result["amount"] == 25.0
|
||||
|
||||
@@ -1,31 +1,76 @@
|
||||
"""scripts_ai (抖音解析/改写/标题) — v1.6.2 起全部免费,不扣积分"""
|
||||
"""scripts_ai 积分扣点单元测试 (#1895 P2 step 2.3)"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
class TestScriptsAiFree:
|
||||
"""三个端点都已移除 @points_gate,不再扣点。"""
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
def test_all_scenes_return_zero_cost(self):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
import packages.middleware.points_gate as _pg_module
|
||||
|
||||
for scene in ("douyin_extract", "ai_rewrite", "ai_title"):
|
||||
assert calculate_points_cost(scene, is_member=False) == 0
|
||||
assert calculate_points_cost(scene, is_member=True) == 0
|
||||
|
||||
def test_no_points_gate_decorators(self):
|
||||
def _make_cu(user_id="u1", is_member=False, member_type=None):
|
||||
cu = MagicMock()
|
||||
cu.user.id = user_id
|
||||
cu.user.is_member = is_member
|
||||
cu.user.member_type = member_type
|
||||
return cu
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _enable_gate(monkeypatch):
|
||||
monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True)
|
||||
yield
|
||||
|
||||
|
||||
class TestScriptsAiPointsGate:
|
||||
"""测试 scripts_ai 三个端点都挂了 @points_gate 并正确扣费。"""
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"scene,endpoint_fn_name",
|
||||
[
|
||||
("douyin_extract", "extract_from_douyin"),
|
||||
("ai_rewrite", "ai_rewrite"),
|
||||
("ai_title", "ai_generate_titles"),
|
||||
],
|
||||
)
|
||||
def test_insufficient_points_raises_402(self, scene, endpoint_fn_name):
|
||||
"""积分不足时抛 402。"""
|
||||
from app.api.routes import scripts_ai
|
||||
from app.schemas.scripts_ai import (
|
||||
AiGenerateTitlesRequest,
|
||||
AiRewriteRequest,
|
||||
ExtractFromDouyinRequest,
|
||||
)
|
||||
|
||||
for fn_name in ("extract_from_douyin", "ai_rewrite", "ai_generate_titles"):
|
||||
fn = getattr(scripts_ai, fn_name)
|
||||
assert not hasattr(fn, "__wrapped__"), f"{fn_name} still has @points_gate"
|
||||
fn = getattr(scripts_ai, endpoint_fn_name)
|
||||
db = MagicMock()
|
||||
cu = _make_cu()
|
||||
if scene == "douyin_extract":
|
||||
req = ExtractFromDouyinRequest(url="https://v.douyin.com/abc/")
|
||||
elif scene == "ai_rewrite":
|
||||
req = AiRewriteRequest(content="测试文案")
|
||||
else:
|
||||
req = AiGenerateTitlesRequest(content="测试文案", count=3)
|
||||
|
||||
def test_module_no_points_imports(self):
|
||||
import inspect
|
||||
with patch("packages.domain.points_service.PointsService") as MockSvc:
|
||||
svc = MagicMock()
|
||||
svc.deduct_points.return_value = {"success": False, "balance": 0}
|
||||
MockSvc.return_value = svc
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
fn(request=req, current_user=cu, db=db)
|
||||
assert ei.value.status_code == 402
|
||||
|
||||
def test_disabled_passthrough_no_user_error(self, monkeypatch):
|
||||
"""关闭时不需要 user/db 也能被装饰器透传(验证 gate 关闭零副作用)。"""
|
||||
from app.api.routes import scripts_ai
|
||||
from app.schemas.scripts_ai import AiRewriteRequest
|
||||
|
||||
src = inspect.getsource(scripts_ai)
|
||||
assert "PointsService" not in src
|
||||
assert "points_gate" not in src
|
||||
assert "calculate_points_cost" not in src
|
||||
monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: False)
|
||||
fn = scripts_ai.ai_rewrite
|
||||
# 不带 db/current_user 也应透传(后续业务逻辑可能报错但不是 401/500 gate 错误)
|
||||
with pytest.raises(Exception) as ei:
|
||||
fn(request=AiRewriteRequest(content="x"), current_user=None, db=None)
|
||||
# 不应是 gate 抛的 401/500
|
||||
assert isinstance(ei.value, AttributeError) or ei.value.status_code not in (401, 500)
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
"""TTS (免费) + voice_clone 预览 (扣点) 单测 (#1895 P2 step 2.1)
|
||||
"""TTS + voice_clone 积分扣点单元测试 (#1895 P2 step 2.1)
|
||||
|
||||
v1.6.2: TTS 合成/预览(ai_voice)已免费,不再扣点;voice_clone 预览(voice_clone_synth)仍保持 1积分/分钟扣点。
|
||||
覆盖 synthesize / voice_clone preview 在积分开关下的扣点、余额不足、失败退费、会员折扣等分支。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -44,10 +44,21 @@ def _make_request(text="你好世界", voice_id="v1", **kw):
|
||||
return r
|
||||
|
||||
|
||||
class TestTtsSynthesizeFree:
|
||||
"""TTS synthesize/preview 已移除手动扣点,不再实例化 PointsService。"""
|
||||
def _est_minutes(chars: int) -> float:
|
||||
return max(1.0, math.ceil(chars / 240))
|
||||
|
||||
def _setup(self, start_synth_raises=None):
|
||||
|
||||
class TestEstimateMinutes:
|
||||
@pytest.mark.parametrize(
|
||||
"chars,expected",
|
||||
[(1, 1.0), (240, 1.0), (241, 2.0), (480, 2.0), (481, 3.0), (1000, 5.0)],
|
||||
)
|
||||
def test_estimate(self, chars, expected):
|
||||
assert _est_minutes(chars) == expected
|
||||
|
||||
|
||||
class TestTtsSynthesizePointsDeduction:
|
||||
def _setup(self, text="你好", deduct_success=True, balance=0, start_synth_raises=None, send_task_raises=None):
|
||||
db = MagicMock()
|
||||
cu = _make_cu()
|
||||
repo = MagicMock()
|
||||
@@ -66,27 +77,43 @@ class TestTtsSynthesizeFree:
|
||||
wf.start_synthesis.side_effect = start_synth_raises
|
||||
vc_repo = MagicMock()
|
||||
vc_repo.get.return_value = None
|
||||
return db, cu, repo, uc, wf, vc_repo, job
|
||||
svc = MagicMock()
|
||||
svc.deduct_points.return_value = {"success": deduct_success, "balance": balance}
|
||||
fake_settings = MagicMock(points_enabled=True)
|
||||
return db, cu, repo, uc, wf, vc_repo, svc, fake_settings, job
|
||||
|
||||
def test_module_has_no_points_imports(self):
|
||||
import inspect
|
||||
|
||||
from app.api.routes import tts as mod
|
||||
|
||||
src = inspect.getsource(mod)
|
||||
assert "PointsService" not in src
|
||||
assert "calculate_points_cost" not in src
|
||||
assert "_points_deducted" not in src
|
||||
assert "import math" not in src
|
||||
|
||||
def test_success_returns_job_without_points(self):
|
||||
db, cu, repo, uc, wf, vc_repo, job = self._setup()
|
||||
def test_insufficient_raises_402(self):
|
||||
db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup(text="你好" * 200, deduct_success=False, balance=0)
|
||||
from app.api.routes.tts import synthesize
|
||||
|
||||
with (
|
||||
patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc),
|
||||
patch("app.api.routes.tts.TTSWorkflowService", return_value=wf),
|
||||
patch("app.api.routes.tts.celery_app.send_task"),
|
||||
patch("app.api.routes.tts.PointsService", return_value=svc),
|
||||
patch("app.api.routes.tts.settings", fs),
|
||||
):
|
||||
with pytest.raises(HTTPException) as ei:
|
||||
synthesize(
|
||||
request=_make_request(text="你好" * 200),
|
||||
authenticated_user=cu,
|
||||
db=db,
|
||||
repository=repo,
|
||||
cosyvoice_service=MagicMock(),
|
||||
voice_clone_repo=vc_repo,
|
||||
)
|
||||
assert ei.value.status_code == 402
|
||||
assert ei.value.detail["code"] == "INSUFFICIENT_POINTS"
|
||||
|
||||
def test_success_deducts_points(self):
|
||||
db, cu, repo, uc, wf, vc_repo, svc, fs, job = self._setup()
|
||||
from app.api.routes.tts import synthesize
|
||||
|
||||
with (
|
||||
patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc),
|
||||
patch("app.api.routes.tts.TTSWorkflowService", return_value=wf),
|
||||
patch("app.api.routes.tts.PointsService", return_value=svc),
|
||||
patch("app.api.routes.tts.celery_app.send_task") as _st,
|
||||
patch("app.api.routes.tts.settings", fs),
|
||||
):
|
||||
resp = synthesize(
|
||||
request=_make_request(text="测试"),
|
||||
@@ -96,17 +123,60 @@ class TestTtsSynthesizeFree:
|
||||
cosyvoice_service=MagicMock(),
|
||||
voice_clone_repo=vc_repo,
|
||||
)
|
||||
svc.deduct_points.assert_called_once()
|
||||
assert resp.job_id == job.id
|
||||
|
||||
def test_ai_voice_cost_zero(self):
|
||||
def test_synthesis_failure_refunds(self):
|
||||
db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup(start_synth_raises=RuntimeError("boom"))
|
||||
from app.api.routes.tts import synthesize
|
||||
|
||||
with (
|
||||
patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc),
|
||||
patch("app.api.routes.tts.TTSWorkflowService", return_value=wf),
|
||||
patch("app.api.routes.tts.PointsService", return_value=svc),
|
||||
patch("app.api.routes.tts.celery_app.send_task"),
|
||||
patch("app.api.routes.tts.settings", fs),
|
||||
):
|
||||
synthesize(
|
||||
request=_make_request(text="测试"),
|
||||
authenticated_user=cu,
|
||||
db=db,
|
||||
repository=repo,
|
||||
cosyvoice_service=MagicMock(),
|
||||
voice_clone_repo=vc_repo,
|
||||
)
|
||||
assert svc.refund_points.called
|
||||
|
||||
def test_celery_send_failure_refunds(self):
|
||||
db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup(send_task_raises=RuntimeError("celery down"))
|
||||
from app.api.routes.tts import synthesize
|
||||
|
||||
with (
|
||||
patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc),
|
||||
patch("app.api.routes.tts.TTSWorkflowService", return_value=wf),
|
||||
patch("app.api.routes.tts.PointsService", return_value=svc),
|
||||
patch("app.api.routes.tts.celery_app.send_task", side_effect=RuntimeError("celery down")),
|
||||
patch("app.api.routes.tts.settings", fs),
|
||||
):
|
||||
synthesize(
|
||||
request=_make_request(text="测试"),
|
||||
authenticated_user=cu,
|
||||
db=db,
|
||||
repository=repo,
|
||||
cosyvoice_service=MagicMock(),
|
||||
voice_clone_repo=vc_repo,
|
||||
)
|
||||
assert svc.refund_points.called
|
||||
|
||||
def test_member_cheaper(self):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
assert calculate_points_cost("ai_voice", is_member=False, duration_minutes=10) == 0
|
||||
cf = calculate_points_cost("ai_voice", is_member=False, duration_minutes=2)
|
||||
cm = calculate_points_cost("ai_voice", is_member=True, member_type="monthly", duration_minutes=2)
|
||||
assert cm < cf
|
||||
|
||||
|
||||
class TestVoiceClonePreviewPoints:
|
||||
"""voice_clone 预览(voice_clone_synth)保持 1 积分/分钟扣点。"""
|
||||
|
||||
def _setup(self, text="你好", deduct_success=True, balance=0, synth_raises=None):
|
||||
db = MagicMock()
|
||||
cu = _make_cu()
|
||||
@@ -195,10 +265,3 @@ class TestVoiceClonePreviewPoints:
|
||||
)
|
||||
svc.deduct_points.assert_called_once()
|
||||
assert resp.audio_url.startswith("http")
|
||||
|
||||
def test_member_cheaper(self):
|
||||
from packages.domain.points_rules import calculate_points_cost
|
||||
|
||||
cf = calculate_points_cost("voice_clone_synth", is_member=False, duration_minutes=2)
|
||||
cm = calculate_points_cost("voice_clone_synth", is_member=True, member_type="monthly", duration_minutes=2)
|
||||
assert cm < cf
|
||||
|
||||
@@ -17,6 +17,7 @@ import pytest
|
||||
from pydantic import ValidationError
|
||||
|
||||
from packages.domain.viral_video import (
|
||||
CREDITS_VIRAL_VIDEO_COST,
|
||||
STAGE_LABELS,
|
||||
FusionLevel,
|
||||
StyleStrength,
|
||||
@@ -128,6 +129,9 @@ class TestViralVideoJobDefaults:
|
||||
assert job.result_video_url == ""
|
||||
assert job.error_msg == ""
|
||||
|
||||
def test_credits_cost_constant(self):
|
||||
assert CREDITS_VIRAL_VIDEO_COST == 50
|
||||
|
||||
|
||||
class TestViralVideoStage:
|
||||
"""阶段枚举测试。"""
|
||||
@@ -555,7 +559,7 @@ class TestPipelineIntegration:
|
||||
mock_review.return_value = {"passed": True, "score": 90}
|
||||
mock_tts.return_value = None # TTS 失败也能走下去(Seedance generate_audio=True 会自己合成音效)
|
||||
mock_tts_upload.return_value = None
|
||||
mock_render.return_value = ("/tmp/video.mp4", {"completion_tokens": 1000000})
|
||||
mock_render.return_value = "/tmp/video.mp4"
|
||||
mock_upload.return_value = "https://oss.example.com/final.mp4"
|
||||
|
||||
result = resume_viral_video_pipeline.run("job-001")
|
||||
@@ -563,4 +567,4 @@ class TestPipelineIntegration:
|
||||
assert result["ok"] is True
|
||||
assert result["video_url"] == "https://oss.example.com/final.mp4"
|
||||
assert job.status == ViralVideoStatus.COMPLETED
|
||||
assert isinstance(job.credits_cost, float)
|
||||
assert job.credits_cost == CREDITS_VIRAL_VIDEO_COST
|
||||
|
||||
@@ -1,38 +1,99 @@
|
||||
"""tests for GET /api/v1/viral-video/models route function (#2159)."""
|
||||
"""爆款视频 DB 模型单元测试(#2039 PR1:DB + migration)。
|
||||
|
||||
验证:
|
||||
- 3 张新表可在内存 SQLite 上创建
|
||||
- 默认值与基本 CRUD 正常
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock, patch
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
Base,
|
||||
ViralVideoJobModel,
|
||||
ViralVideoPromptTemplateModel,
|
||||
ViralVideoStyleTemplateModel,
|
||||
)
|
||||
|
||||
|
||||
class TestModelsRoute:
|
||||
def _call(self):
|
||||
from apps.api.app.api.routes import viral_video as routes
|
||||
def _make_session():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
return sessionmaker(bind=engine)()
|
||||
|
||||
return routes.list_available_models()
|
||||
|
||||
def test_without_dashscope_hides_wan(self):
|
||||
with patch("apps.api.app.api.routes.viral_video.get_dashscope_client", return_value=None):
|
||||
result = self._call()
|
||||
assert "models" in result
|
||||
keys = {m["key"] for m in result["models"]}
|
||||
assert "seedance-2.5" in keys
|
||||
assert "wan-3.0" not in keys
|
||||
for m in result["models"]:
|
||||
if m["key"].startswith("seedance"):
|
||||
assert m["supports_audio"] is True
|
||||
class TestViralVideoJobModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(
|
||||
id="job-001",
|
||||
user_id="user-001",
|
||||
images=["https://img.com/1.jpg"],
|
||||
industry="美妆",
|
||||
duration=60,
|
||||
)
|
||||
session.add(job)
|
||||
session.commit()
|
||||
|
||||
def test_with_dashscope_includes_wan(self):
|
||||
with patch("apps.api.app.api.routes.viral_video.get_dashscope_client", return_value=MagicMock()):
|
||||
result = self._call()
|
||||
keys = {m["key"] for m in result["models"]}
|
||||
assert "wan-3.0" in keys
|
||||
wan = next(m for m in result["models"] if m["key"] == "wan-3.0")
|
||||
assert wan["billing_mode"] == "per_second"
|
||||
fetched = session.query(ViralVideoJobModel).filter_by(id="job-001").one()
|
||||
assert fetched.user_id == "user-001"
|
||||
assert fetched.images == ["https://img.com/1.jpg"]
|
||||
assert fetched.industry == "美妆"
|
||||
assert fetched.duration == 60
|
||||
|
||||
def test_exactly_one_default(self):
|
||||
with patch("apps.api.app.api.routes.viral_video.get_dashscope_client", return_value=None):
|
||||
result = self._call()
|
||||
defaults = [m for m in result["models"] if m["is_default"]]
|
||||
assert len(defaults) == 1
|
||||
assert defaults[0]["key"] == "seedance-2.5"
|
||||
def test_default_values(self):
|
||||
session = _make_session()
|
||||
job = ViralVideoJobModel(id="job-002", user_id="user-002")
|
||||
session.add(job)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoJobModel, "job-002")
|
||||
assert fetched.images == []
|
||||
assert fetched.fusion_level == "ai_polish"
|
||||
assert fetched.style_strength == "medium"
|
||||
assert fetched.status == "pending"
|
||||
assert fetched.credits_cost == 0
|
||||
assert fetched.retry_count == 0
|
||||
assert fetched.style_guide is None
|
||||
assert fetched.intent_result is None
|
||||
|
||||
|
||||
class TestViralVideoStyleTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoStyleTemplateModel(
|
||||
id="tpl-001",
|
||||
name="快节奏",
|
||||
style_config={"cut_speed": "fast"},
|
||||
sort_order=1,
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoStyleTemplateModel, "tpl-001")
|
||||
assert fetched.name == "快节奏"
|
||||
assert fetched.style_config == {"cut_speed": "fast"}
|
||||
assert fetched.sort_order == 1
|
||||
|
||||
|
||||
class TestViralVideoPromptTemplateModel:
|
||||
def test_create_and_get(self):
|
||||
session = _make_session()
|
||||
tpl = ViralVideoPromptTemplateModel(
|
||||
id="pt-001",
|
||||
prompt_type="image_analysis",
|
||||
name="图片分析模板",
|
||||
content="请分析图片:{image_url}",
|
||||
variables=["image_url"],
|
||||
)
|
||||
session.add(tpl)
|
||||
session.commit()
|
||||
|
||||
fetched = session.get(ViralVideoPromptTemplateModel, "pt-001")
|
||||
assert fetched.prompt_type == "image_analysis"
|
||||
assert fetched.content == "请分析图片:{image_url}"
|
||||
assert fetched.variables == ["image_url"]
|
||||
assert fetched.version == 1
|
||||
assert fetched.is_active is True
|
||||
|
||||
@@ -212,13 +212,10 @@ class TestCallVideoGeneration:
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.video_generation.return_value = {
|
||||
"video_path": str(out),
|
||||
"usage": {"completion_tokens": 1000000},
|
||||
}
|
||||
mock_client.video_generation.return_value = str(out)
|
||||
mock_get.return_value = mock_client
|
||||
result = call_video_generation(prompt="测试", image_url="https://img/x.jpg", duration=5, ratio="9:16")
|
||||
assert result is not None and result["video_path"] == str(out)
|
||||
assert result == str(out)
|
||||
mock_client.video_generation.assert_called_once()
|
||||
kwargs = mock_client.video_generation.call_args.kwargs
|
||||
assert kwargs["prompt"] == "测试"
|
||||
@@ -241,10 +238,7 @@ class TestCallVideoGenerationV16:
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.video_generation.return_value = {
|
||||
"video_path": str(out),
|
||||
"usage": {"completion_tokens": 1500000},
|
||||
}
|
||||
mock_client.video_generation.return_value = str(out)
|
||||
mock_get.return_value = mock_client
|
||||
result = call_video_generation(
|
||||
prompt="测试",
|
||||
@@ -257,7 +251,7 @@ class TestCallVideoGenerationV16:
|
||||
generate_audio=True,
|
||||
model="doubao-seedance-2-5-260628",
|
||||
)
|
||||
assert result is not None and result["video_path"] == str(out)
|
||||
assert result == str(out)
|
||||
kwargs = mock_client.video_generation.call_args.kwargs
|
||||
# 图生视频也必须传 ratio(避免首帧方图导致默认输出 1:1)
|
||||
assert kwargs.get("ratio") == "9:16", f"ratio 应透传,got {kwargs.get('ratio')!r}"
|
||||
@@ -276,7 +270,7 @@ class TestCallVideoGenerationV16:
|
||||
with patch("packages.shared.ai_service.get_doubao_client") as mock_get:
|
||||
mock_client = MagicMock()
|
||||
mock_client.is_available = True
|
||||
mock_client.video_generation.return_value = {"video_path": str(out), "usage": None}
|
||||
mock_client.video_generation.return_value = str(out)
|
||||
mock_get.return_value = mock_client
|
||||
call_video_generation(prompt="测试", duration=10, ratio="16:9")
|
||||
kwargs = mock_client.video_generation.call_args.kwargs
|
||||
@@ -313,133 +307,3 @@ class TestResumeReadsImageAnalysis:
|
||||
# resume 本身应该调用 _run_render_pipeline
|
||||
resume_src = inspect.getsource(vv.resume_viral_video_pipeline)
|
||||
assert "_run_render_pipeline" in resume_src
|
||||
|
||||
|
||||
# ============ #2157 _replace_henjin_everywhere 递归替换单测 ============
|
||||
|
||||
|
||||
class TestReplaceHenjinEverywhere:
|
||||
"""覆盖 #2157 P1:递归替换 copy_result 中所有层级的'很近'→'最近'。"""
|
||||
|
||||
def _import(self):
|
||||
from apps.worker.worker_app.tasks.viral_video import _replace_henjin_everywhere
|
||||
|
||||
return _replace_henjin_everywhere
|
||||
|
||||
def test_plain_string_no_henjin(self):
|
||||
fn = self._import()
|
||||
assert fn("最近好物推荐") == "最近好物推荐"
|
||||
assert fn("") == ""
|
||||
assert fn(None) is None
|
||||
assert fn(123) == 123
|
||||
|
||||
def test_string_with_henjin(self):
|
||||
fn = self._import()
|
||||
assert fn("很近是不是总觉得颈肩发僵") == "最近是不是总觉得颈肩发僵"
|
||||
# 多次出现
|
||||
assert fn("很近很近都很近") == "最近最近都最近"
|
||||
|
||||
def test_list_recursive(self):
|
||||
fn = self._import()
|
||||
out = fn(["很近a", "b", ["很近c", "d"]])
|
||||
assert out == ["最近a", "b", ["最近c", "d"]]
|
||||
|
||||
def test_dict_recursive_nested(self):
|
||||
fn = self._import()
|
||||
obj = {
|
||||
"overview": {"theme": "很近颈肩", "title": "x"},
|
||||
"scene_and_lighting": "很近才好用",
|
||||
"voiceover_script": "很近是不是",
|
||||
"final_copy": "很近好物",
|
||||
"shots": [
|
||||
{"scene_and_dialogue": "很近第一镜", "action_details": "很近动作", "audio_bgm": "很近音乐"},
|
||||
{"nested": {"deep": "很近深层"}},
|
||||
],
|
||||
"int_field": 42,
|
||||
}
|
||||
import json
|
||||
|
||||
out = fn(obj)
|
||||
assert "很近" not in json.dumps(out, ensure_ascii=False)
|
||||
assert out["overview"]["theme"] == "最近颈肩"
|
||||
assert out["shots"][0]["scene_and_dialogue"] == "最近第一镜"
|
||||
assert out["shots"][1]["nested"]["deep"] == "最近深层"
|
||||
assert out["int_field"] == 42
|
||||
|
||||
|
||||
class TestMarkFailedAndNotifySessionFallback:
|
||||
"""#2157 P1:_mark_failed_and_notify 在原session失效时fallback到新SessionLocal。"""
|
||||
|
||||
def test_fallback_to_new_session_when_original_save_raises(self, tmp_path):
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
job = MagicMock()
|
||||
job.is_terminal = False
|
||||
job.mark_failed = MagicMock()
|
||||
|
||||
# 原 session 保存抛异常
|
||||
orig_session = MagicMock()
|
||||
orig_repo = MagicMock()
|
||||
|
||||
def _raise(*a, **kw):
|
||||
raise RuntimeError("session in rollback")
|
||||
|
||||
# 第一次调用_save_job抛异常,触发fallback
|
||||
with patch.object(vv, "_save_job", side_effect=_raise):
|
||||
fake_ssn = MagicMock()
|
||||
fake_repo = MagicMock()
|
||||
fake_job_in_db = MagicMock()
|
||||
fake_job_in_db.is_terminal = False
|
||||
fake_repo.get.return_value = fake_job_in_db
|
||||
with patch.object(vv, "SessionLocal", return_value=fake_ssn):
|
||||
with patch.object(vv, "SQLAlchemyViralVideoJobRepository", return_value=fake_repo):
|
||||
with patch.object(vv, "_emit_progress") as mock_emit:
|
||||
vv._mark_failed_and_notify("job-1", orig_session, orig_repo, job, "boom", stage="render")
|
||||
# 原session上mark_failed被调用过
|
||||
job.mark_failed.assert_called()
|
||||
# fallback路径:新session上repo.get(job-1)被调用,且新job被mark_failed并commit
|
||||
fake_repo.get.assert_called_with("job-1")
|
||||
fake_job_in_db.mark_failed.assert_called_with("boom")
|
||||
fake_repo.update.assert_called_with(fake_job_in_db)
|
||||
fake_ssn.commit.assert_called()
|
||||
fake_ssn.close.assert_called()
|
||||
mock_emit.assert_called_once()
|
||||
|
||||
def test_original_session_happy_path_no_fallback(self):
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
job = MagicMock()
|
||||
job.is_terminal = False
|
||||
session = MagicMock()
|
||||
repo = MagicMock()
|
||||
with patch.object(vv, "_save_job") as mock_save:
|
||||
with patch.object(vv, "_emit_progress") as mock_emit:
|
||||
# 不mock SessionLocal,如果fallback被错误触发会抛AttributeError
|
||||
vv._mark_failed_and_notify("job-2", session, repo, job, "err", stage="copy")
|
||||
job.mark_failed.assert_called_with("err")
|
||||
mock_save.assert_called()
|
||||
mock_emit.assert_called_once()
|
||||
|
||||
def test_terminal_job_not_marked(self):
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
from apps.worker.worker_app.tasks import viral_video as vv
|
||||
|
||||
job = MagicMock()
|
||||
job.is_terminal = True # 已终态
|
||||
session = MagicMock()
|
||||
repo = MagicMock()
|
||||
with patch.object(vv, "_save_job") as mock_save:
|
||||
with patch.object(vv, "_emit_progress"):
|
||||
fake_ssn = MagicMock()
|
||||
with patch.object(vv, "SessionLocal", return_value=fake_ssn):
|
||||
with patch.object(vv, "SQLAlchemyViralVideoJobRepository") as mock_repo_cls:
|
||||
vv._mark_failed_and_notify("job-3", session, repo, job, "x")
|
||||
# 终态job不调用mark_failed
|
||||
job.mark_failed.assert_not_called()
|
||||
# 且因 job 已终态,_save_job 也不应被调用(marked=False 才fallback;但此处 job 非 None 且 is_terminal=True,marked 保持 False 进入 fallback)
|
||||
# fallback路径会重新打开session,get到的job也是终态,不会update
|
||||
|
||||
@@ -14,8 +14,6 @@ from __future__ import annotations
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
def _auth_user(uid: str = "u1"):
|
||||
return SimpleNamespace(user=SimpleNamespace(id=uid))
|
||||
@@ -62,10 +60,7 @@ def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending
|
||||
"voice_mode": "global",
|
||||
"video_ratio": "9:16",
|
||||
"video_model": "",
|
||||
"video_resolution": "720p",
|
||||
"credits_prepaid": 0.0,
|
||||
"credits_transaction_id": "",
|
||||
"credits_cost": 0.0,
|
||||
"credits_cost": 0,
|
||||
"current_stage": "",
|
||||
"phase_message": "",
|
||||
"updated_at": None,
|
||||
@@ -138,163 +133,6 @@ class TestRetryViralVideo:
|
||||
mock_send.assert_called_once_with("worker.run_viral_video_pipeline", args=["job-retry"])
|
||||
assert resp.id == "job-retry"
|
||||
|
||||
def test_retry_without_body_keeps_original_params(self):
|
||||
"""不传 body 时,保持原参数且不调用积分服务。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
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,
|
||||
)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch("app.config.settings") as mock_settings,
|
||||
patch("packages.domain.points_service.PointsService") as MockSvc,
|
||||
patch.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
# request=None (未传 body)
|
||||
resp = vv_mod.retry_viral_video_job("job-retry2", None, authenticated_user=user, session=session)
|
||||
|
||||
MockSvc.assert_not_called()
|
||||
assert resp.id == "job-retry2"
|
||||
assert job.status == ViralVideoStatus.PENDING
|
||||
assert job.duration == 15 # 参数不变
|
||||
|
||||
def test_retry_insufficient_points_raises_402(self):
|
||||
"""参数变更导致新预估更高且余额不足时,抛 402 阻止重试。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import RetryViralVideoRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
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,
|
||||
)
|
||||
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")
|
||||
|
||||
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.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
vv_mod.retry_viral_video_job("job-retry3a", req, authenticated_user=user, session=session)
|
||||
assert exc.value.status_code == 402
|
||||
fake_svc.deduct_viral_video.assert_called_once()
|
||||
|
||||
def test_retry_higher_estimation_calls_deduct_delta(self):
|
||||
"""参数变更新预估更高时调用 deduct_viral_video 补扣差额,并更新 job 参数。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import RetryViralVideoRequest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
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,
|
||||
)
|
||||
# 用 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")
|
||||
|
||||
new_est = 15.0
|
||||
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=(new_est, {})),
|
||||
patch.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
resp = vv_mod.retry_viral_video_job("job-retry3b", req, authenticated_user=user, session=session)
|
||||
|
||||
assert job.duration == 30
|
||||
assert job.video_resolution == "1080p"
|
||||
assert job.video_ratio == "16:9"
|
||||
# 补扣差额 = 15-5 = 10
|
||||
fake_svc.deduct_viral_video.assert_called_once()
|
||||
call_args = fake_svc.deduct_viral_video.call_args
|
||||
assert call_args.args[1] == 10.0 # credits 是位置参数
|
||||
assert resp.id == "job-retry3b"
|
||||
|
||||
def test_retry_lower_estimation_calls_refund_delta(self):
|
||||
"""参数变更新预估更低时,调用 refund_points 退还差额,并更新 job 参数。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import RetryViralVideoRequest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
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.credits_prepaid = 10.0
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
|
||||
fake_svc = MagicMock()
|
||||
fake_svc.refund_points.return_value = {"success": True}
|
||||
req = RetryViralVideoRequest(duration=5, video_resolution="480p", video_ratio="9:16")
|
||||
|
||||
new_est = 3.0
|
||||
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=(270, 480)),
|
||||
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
|
||||
vv_mod.retry_viral_video_job("job-retry4", req, authenticated_user=user, session=session)
|
||||
|
||||
assert job.duration == 5
|
||||
assert job.video_resolution == "480p"
|
||||
assert job.video_ratio == "9:16"
|
||||
fake_svc.refund_points.assert_called_once()
|
||||
call_args = fake_svc.refund_points.call_args
|
||||
# 退差额 = 10-3 = 7
|
||||
assert call_args.kwargs["amount"] == 7.0
|
||||
|
||||
|
||||
# ── confirm-intent ──────────────────────────────────────────────────────
|
||||
|
||||
@@ -561,306 +399,3 @@ class TestConfirmCopy:
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
vv_mod.confirm_copy("job-cc2", ConfirmCopyRequest(), authenticated_user=user, session=session)
|
||||
assert exc.value.status_code == 409
|
||||
|
||||
|
||||
# ── confirm-copy 积分预扣 + estimate-credits 端点 (#2151) ──────────────
|
||||
|
||||
|
||||
class TestConfirmCopyPointsDeduction:
|
||||
"""confirm_copy 中积分预扣分支(points_enabled=True)。"""
|
||||
|
||||
def test_points_enabled_deducts_successfully(self):
|
||||
"""points_enabled=True + 未预付 → 计算预估积分 → deduct_viral_video → 写入 credits_prepaid。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import ConfirmCopyRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-pay", user_id="u1", status=ViralVideoStatus.COPY_GENERATED)
|
||||
# 默认 credits_prepaid=0, credits_cost=0 → 触发预扣
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = ConfirmCopyRequest(edited_copy="改好的文案")
|
||||
|
||||
mock_svc = MagicMock()
|
||||
mock_svc.deduct_viral_video.return_value = {"success": True, "balance": 100.0, "transaction_id": "txn-1"}
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch.object(vv_mod, "_settings", None, create=True), # ensure not cached
|
||||
patch("app.config.settings") as mock_settings,
|
||||
patch("packages.domain.points_rules.calculate_viral_video_credits", return_value=5.2),
|
||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(720, 1280)),
|
||||
patch("packages.domain.points_service.PointsService", return_value=mock_svc),
|
||||
patch.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
resp = vv_mod.confirm_copy("job-pay", req, authenticated_user=user, session=session)
|
||||
|
||||
# deduct_viral_video 被调用
|
||||
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[2] == "job-pay" # job_id
|
||||
# credits_prepaid / credits_transaction_id 被写入
|
||||
assert job.credits_prepaid == 5.2
|
||||
assert job.credits_transaction_id == "txn-1"
|
||||
assert resp.id == "job-pay"
|
||||
# resume + repo.update 至少调用过(其中一次是 credits 字段更新,一次是 resume 后)
|
||||
job.resume_from_copy_generated.assert_called_once_with(edited_copy="改好的文案")
|
||||
|
||||
def test_points_enabled_insufficient_balance_raises_402(self):
|
||||
"""余额不足(deduct_viral_video 返回 success=False)→ HTTP 402。"""
|
||||
import pytest
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import ConfirmCopyRequest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-402", user_id="u1", status=ViralVideoStatus.COPY_GENERATED)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = ConfirmCopyRequest()
|
||||
|
||||
mock_svc = MagicMock()
|
||||
mock_svc.deduct_viral_video.return_value = {"success": False, "balance": 1.5, "transaction_id": None}
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch("app.config.settings") as mock_settings,
|
||||
patch("packages.domain.points_rules.calculate_viral_video_credits", return_value=10.0),
|
||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(720, 1280)),
|
||||
patch("packages.domain.points_service.PointsService", return_value=mock_svc),
|
||||
patch.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
with pytest.raises(HTTPException) as exc:
|
||||
vv_mod.confirm_copy("job-402", req, authenticated_user=user, session=session)
|
||||
|
||||
assert exc.value.status_code == 402
|
||||
detail = exc.value.detail
|
||||
assert detail["code"] == "INSUFFICIENT_POINTS"
|
||||
assert detail["required"] == 10.0
|
||||
assert detail["balance"] == 1.5
|
||||
# 预扣失败不应调用 resume 或 send_task
|
||||
job.resume_from_copy_generated.assert_not_called()
|
||||
|
||||
def test_already_paid_skips_deduction(self):
|
||||
"""credits_prepaid>0(已经扣过费/重试场景) → 跳过预扣,不调用 PointsService。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import ConfirmCopyRequest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(
|
||||
job_id="job-paid",
|
||||
user_id="u1",
|
||||
status=ViralVideoStatus.COPY_GENERATED,
|
||||
credits_prepaid=8.5,
|
||||
credits_transaction_id="txn-old",
|
||||
)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = ConfirmCopyRequest(edited_copy="继续")
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch("app.config.settings") as mock_settings,
|
||||
patch("packages.domain.points_service.PointsService") as MockSvc,
|
||||
patch.object(vv_mod.celery_app, "send_task") as mock_send,
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
resp = vv_mod.confirm_copy("job-paid", req, authenticated_user=user, session=session)
|
||||
|
||||
# PointsService 不应被实例化(没预扣)
|
||||
MockSvc.assert_not_called()
|
||||
job.resume_from_copy_generated.assert_called_once_with(edited_copy="继续")
|
||||
mock_send.assert_called_once_with("worker.run_viral_video_render", args=["job-paid"])
|
||||
assert resp.id == "job-paid"
|
||||
# credits_prepaid 保持不变
|
||||
assert job.credits_prepaid == 8.5
|
||||
|
||||
def test_already_paid_via_credits_cost_skips_deduction(self):
|
||||
"""credits_cost>0 也算已付费(兼容旧字段),跳过预扣。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import ConfirmCopyRequest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(
|
||||
job_id="job-paid2",
|
||||
user_id="u1",
|
||||
status=ViralVideoStatus.COPY_GENERATED,
|
||||
credits_prepaid=0,
|
||||
credits_cost=7.0,
|
||||
)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = ConfirmCopyRequest()
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch("app.config.settings") as mock_settings,
|
||||
patch("packages.domain.points_service.PointsService") as MockSvc,
|
||||
patch.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = True
|
||||
vv_mod.confirm_copy("job-paid2", req, authenticated_user=user, session=session)
|
||||
|
||||
MockSvc.assert_not_called()
|
||||
|
||||
def test_points_disabled_skips_deduction(self):
|
||||
"""points_enabled=False 时不进入预扣逻辑,保持原流程。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import ConfirmCopyRequest
|
||||
|
||||
from packages.domain.viral_video import ViralVideoStatus
|
||||
|
||||
user = _auth_user("u1")
|
||||
session = MagicMock()
|
||||
job = _make_job(job_id="job-free", user_id="u1", status=ViralVideoStatus.COPY_GENERATED)
|
||||
repo = MagicMock()
|
||||
repo.get.return_value = job
|
||||
req = ConfirmCopyRequest()
|
||||
|
||||
with (
|
||||
patch.object(vv_mod, "_get_job_repo", return_value=repo),
|
||||
patch("app.config.settings") as mock_settings,
|
||||
patch("packages.domain.points_service.PointsService") as MockSvc,
|
||||
patch.object(vv_mod.celery_app, "send_task"),
|
||||
):
|
||||
mock_settings.points_enabled = False
|
||||
vv_mod.confirm_copy("job-free", req, authenticated_user=user, session=session)
|
||||
|
||||
MockSvc.assert_not_called()
|
||||
job.resume_from_copy_generated.assert_called_once()
|
||||
|
||||
|
||||
class TestEstimateCredits:
|
||||
"""POST /estimate-credits: 纯计算预估积分。"""
|
||||
|
||||
def test_estimate_returns_float_with_breakdown(self):
|
||||
"""正常参数应返回 estimated_credits(float, >0, 两位小数) + formula_breakdown。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import EstimateCreditsRequest
|
||||
|
||||
req = EstimateCreditsRequest(model="seedance-2.5", resolution="720p", ratio="9:16", duration=15)
|
||||
user = _auth_user("u1")
|
||||
|
||||
resp = vv_mod.estimate_credits(req, authenticated_user=user)
|
||||
assert isinstance(resp.estimated_credits, float)
|
||||
assert resp.estimated_credits > 0
|
||||
assert round(resp.estimated_credits, 2) == resp.estimated_credits
|
||||
# formula_breakdown 必须返回并包含全部字段
|
||||
bd = resp.formula_breakdown
|
||||
assert bd.tokens > 0
|
||||
assert bd.video_cost >= 0
|
||||
assert bd.fixed_cost > 0
|
||||
assert bd.profit_multiplier == 1.3
|
||||
assert bd.model_price > 0
|
||||
assert bd.width > 0
|
||||
assert bd.height > 0
|
||||
assert bd.fps > 0
|
||||
|
||||
def test_estimate_uses_dimensions_resolver_and_with_breakdown(self):
|
||||
"""estimate_credits 调用 resolve_video_dimensions 与 calculate_viral_video_credits_with_breakdown。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import EstimateCreditsRequest
|
||||
|
||||
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,
|
||||
}
|
||||
|
||||
with (
|
||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(1920, 1080)) as mock_res,
|
||||
patch(
|
||||
"packages.domain.points_rules.calculate_viral_video_credits_with_breakdown",
|
||||
return_value=(8.88, fake_bd),
|
||||
) as mock_calc,
|
||||
):
|
||||
resp = vv_mod.estimate_credits(req, authenticated_user=user)
|
||||
|
||||
mock_res.assert_called_once_with("1080p", "16:9")
|
||||
mock_calc.assert_called_once()
|
||||
args, kwargs = mock_calc.call_args
|
||||
assert args[0] == 20
|
||||
assert args[1] == 1920
|
||||
assert args[2] == 1080
|
||||
assert args[3] == "seedance-2.5"
|
||||
assert resp.estimated_credits == 8.88
|
||||
assert resp.formula_breakdown.width == 1920
|
||||
assert resp.formula_breakdown.height == 1080
|
||||
assert resp.formula_breakdown.model_price == 70.0
|
||||
|
||||
def test_estimate_empty_model_defaults_to_seedance_2_5(self):
|
||||
"""model 为空字符串时,传入 calculate 的 model 参数应为 'seedance-2.5'。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import EstimateCreditsRequest
|
||||
|
||||
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,
|
||||
}
|
||||
|
||||
with (
|
||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(720, 1280)),
|
||||
patch(
|
||||
"packages.domain.points_rules.calculate_viral_video_credits_with_breakdown",
|
||||
return_value=(3.5, fake_bd),
|
||||
) as mock_calc,
|
||||
):
|
||||
resp = vv_mod.estimate_credits(req, authenticated_user=user)
|
||||
|
||||
args, kwargs = mock_calc.call_args
|
||||
assert args[3] == "seedance-2.5"
|
||||
assert resp.estimated_credits == 3.5
|
||||
assert resp.formula_breakdown.height == 1280
|
||||
|
||||
def test_estimate_accepts_video_model_alias(self):
|
||||
"""前端传 video_model/video_resolution/video_ratio(别名)也应被正确解析。"""
|
||||
from app.api.routes import viral_video as vv_mod
|
||||
from app.schemas.viral_video import EstimateCreditsRequest
|
||||
|
||||
req = EstimateCreditsRequest.model_validate(
|
||||
{"video_model": "seedance-2.0", "video_resolution": "480p", "video_ratio": "1:1", "duration": 5}
|
||||
)
|
||||
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,
|
||||
}
|
||||
|
||||
with (
|
||||
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(480, 480)) as mock_res,
|
||||
patch(
|
||||
"packages.domain.points_rules.calculate_viral_video_credits_with_breakdown",
|
||||
return_value=(1.0, fake_bd),
|
||||
) as mock_calc,
|
||||
):
|
||||
vv_mod.estimate_credits(req, authenticated_user=user)
|
||||
|
||||
mock_res.assert_called_once_with("480p", "1:1")
|
||||
args, kwargs = mock_calc.call_args
|
||||
assert args[0] == 5
|
||||
assert args[1] == 480
|
||||
assert args[2] == 480
|
||||
assert args[3] == "seedance-2.0"
|
||||
|
||||
Reference in New Issue
Block a user