Compare commits

..

1 Commits

Author SHA1 Message Date
Coze Agent c5e0cbba6a fix(credits): 爆款视频固定显示50积分,6项功能积分展示已由全局开关隐藏
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 1s
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 1s
CI/CD Pipeline / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 47s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m45s
PR Automation / Auto Approve on CI Green (pull_request) Failing after 2m5s
CI/CD Pipeline / PR Build Web Image (pull_request) Successful in 2m22s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 2m36s
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 3m45s
CI/CD Pipeline / Frontend Unit Tests (pull_request) Successful in 4m15s
AI Code Review / AI Code Review (pull_request) Successful in 7m27s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been skipped
CI/CD Pipeline / Validate - Style (pull_request) Successful in 15m14s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 20m42s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 43m31s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
2026-10-02 20:23:10 +08:00
61 changed files with 1672 additions and 6212 deletions
-14
View File
@@ -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),
+94 -4
View File
@@ -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}",
+32 -18
View File
@@ -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,
)
+4
View File
@@ -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),
+24 -4
View File
@@ -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])
+76
View File
@@ -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
+3 -203
View File
@@ -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}")
+10 -10
View File
@@ -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)",
+1 -58
View File
@@ -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(
-25
View File
@@ -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 辅助函数(后端新接口上线后可替换) ── */
/**
-19
View File
@@ -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) */
+2 -22
View File
@@ -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,
+2 -6
View File
@@ -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
}
+100 -156
View File
@@ -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;
+202 -636
View File
@@ -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}
>
+30 -254
View File
@@ -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:
+7 -10
View File
@@ -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
View File
@@ -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
View File
@@ -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"
-5
View File
@@ -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))
+3 -3
View File
@@ -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
View File
@@ -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
+129 -98
View File
@@ -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(),
}
+3 -4
View File
@@ -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
+13
View File
@@ -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
View File
@@ -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
# ── 单例 ─────────────────────────────────────────────────────────────────────
+9 -31
View File
@@ -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 {}
-344
View File
@@ -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
-526
View File
@@ -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
+28 -19
View File
@@ -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 = {}
+1 -3
View File
@@ -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)
+37 -14
View File
@@ -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
+10 -187
View File
@@ -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
-198
View File
@@ -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
+14 -46
View File
@@ -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
-178
View File
@@ -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
+19 -16
View File
@@ -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"
+43 -20
View File
@@ -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"
+53 -18
View File
@@ -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"
-376
View File
@@ -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"]
+189 -60
View File
@@ -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)
+13 -20
View File
@@ -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)
+50 -72
View File
@@ -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
View File
@@ -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)
+18 -181
View File
@@ -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
+63 -18
View File
@@ -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)
+94 -31
View File
@@ -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
+6 -2
View File
@@ -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
+90 -29
View File
@@ -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
+5 -141
View File
@@ -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
+1 -466
View File
@@ -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"