feat(viral-video): 积分预估返回formula_breakdown + retry改参多退少补 + POINTS_SCENES注册viral_video (#2153)
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (push) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 2s
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 (push) Successful in 2s
CI/CD Pipeline / PR Build API Image (push) Has been skipped
CI/CD Pipeline / PR Build Web Image (push) Has been skipped
CI/CD Pipeline / PR Build Worker Image (push) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Validate - Style (pull_request) Has been skipped
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been skipped
CI/CD Pipeline / Validate - Security (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / Check push changed paths (push) Successful in 18s
CI/CD Pipeline / Frontend Lint (push) 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 / 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 / PR Build API Image (pull_request) Successful in 1m7s
CI/CD Pipeline / Build Staging API Image (push) Successful in 1m52s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (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
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 2m34s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m11s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been skipped
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 3m18s
CI/CD Pipeline / Build Staging Worker Image (push) Successful in 2m15s
CI/CD Pipeline / CI Gate (pull_request) Successful in 2s
CI/CD Pipeline / Build Staging Web Image (push) Successful in 2m45s
CI/CD Pipeline / Retag skipped Staging Web Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging API Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (push) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (push) Successful in 4m13s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Successful in 1m27s
CI/CD Pipeline / Validate - Style (push) Successful in 6m8s
CI/CD Pipeline / Integration Tests (push) Successful in 6m25s
AI Code Review / AI Code Review (pull_request) Successful in 7m23s
CI/CD Pipeline / Validate - Python (mypy + alembic) (push) Successful in 7m51s
CI/CD Pipeline / ACR Image Cleanup (push) Successful in 3m26s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 4m13s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 4m50s
CI/CD Pipeline / Validate - Security (push) Successful in 13m4s
CI/CD Pipeline / Unit Tests (push) Successful in 15m56s
CI/CD Pipeline / Build Production API Image (push) Has been skipped
CI/CD Pipeline / Build Production Web Image (push) Has been skipped
CI/CD Pipeline / CI Gate (push) Has been skipped
CI/CD Pipeline / Build Production Worker Image (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Canary Release to Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped

Co-authored-by: xiaoxia <dev@xiaoxiajianji.com>
Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
This commit was merged in pull request #2153.
This commit is contained in:
2026-10-03 01:00:22 +08:00
committed by auto-approve-bot
parent 7800ff4c3d
commit bd617128ce
6 changed files with 558 additions and 48 deletions
+128 -8
View File
@@ -32,9 +32,11 @@ from app.schemas.viral_video import (
ConfirmCopyRequest,
ConfirmIntentRequest,
CreateViralVideoRequest,
CreditsFormulaBreakdown,
EstimateCreditsRequest,
EstimateCreditsResponse,
GenerateCopyRequest,
RetryViralVideoRequest,
StyleTemplateListResponse,
StyleTemplateResponse,
ViralVideoHistoryResponse,
@@ -391,14 +393,28 @@ def estimate_credits(
request: EstimateCreditsRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
) -> EstimateCreditsResponse:
"""爆款视频积分预估(纯计算,不扣费、不创建任务)。"""
from packages.domain.points_rules import calculate_viral_video_credits, resolve_video_dimensions
"""爆款视频积分预估(纯计算,不扣费、不创建任务)。
w, h = resolve_video_dimensions(request.resolution, request.ratio)
credits = calculate_viral_video_credits(
request.duration, w, h, request.model or "seedance-2.5"
返回 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,
)
return EstimateCreditsResponse(estimated_credits=credits)
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)
@@ -454,10 +470,17 @@ 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 任务强制重置后重试)。"""
"""重试失败的爆款视频任务(也支持对僵尸/超时 running 任务强制重置后重试)。
可选 body (RetryViralVideoRequest):若传入新的 duration/video_resolution/video_ratio/
video_model,会重新预估积分并与原 credits_prepaid 做差额多退少补(不足抛 402 阻止重试);
不传 body 或参数无变化时,保持原参数、原预扣金额不变,仅重置状态并入队。
credits_prepaid 为 0 的老任务首次重试会走预扣流程(与 confirm-copy 一致)。
"""
from datetime import datetime, timezone
repo = _get_job_repo(session)
@@ -478,6 +501,100 @@ 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
@@ -492,7 +609,10 @@ 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", job.id, job.retry_count, is_stale_running)
logger.info(
"[爆款视频] 重试入队: job_id=%s retry_count=%d stale=%s params_changed=%s",
job.id, job.retry_count, is_stale_running, param_changed,
)
except Exception as e:
logger.error("[爆款视频] 重试入队失败: %s", e, exc_info=True)
job.mark_failed(f"重试入队失败: {e}")
+37 -4
View File
@@ -260,18 +260,51 @@ class AnalyzeStyleResponse(BaseModel):
class EstimateCreditsRequest(BaseModel):
"""爆款视频积分预估请求。"""
"""爆款视频积分预估请求。
model: str = ""
resolution: str = "720p"
ratio: str = "9:16"
前端可传 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 --
+88 -24
View File
@@ -86,6 +86,61 @@ def _infer_resolution_key(height: int) -> str:
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 字段,便于前端展示计费明细。
"""
prefix = _match_model_prefix(model)
res_key = _infer_resolution_key(int(height or 720))
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)
w = max(1, int(width or 1))
h = max(1, int(height or 1))
effective_fps = int(fps or VIRAL_VIDEO_FPS)
if actual_tokens is not None and actual_tokens > 0:
tokens = float(actual_tokens)
else:
dur = max(1, int(duration_seconds or 15))
tokens = dur * w * h * effective_fps / 1024.0
video_cost = tokens / 1_000_000.0 * float(price)
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),
"width": int(w),
"height": int(h),
"fps": int(effective_fps),
}
return credits, breakdown
def calculate_viral_video_credits(
duration_seconds: int,
width: int,
@@ -95,37 +150,33 @@ def calculate_viral_video_credits(
actual_tokens: int | None = None,
fps: int = VIRAL_VIDEO_FPS,
) -> float:
"""计算爆款视频所需积分(1 积分 = 1 元)。
"""计算爆款视频所需积分(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) * 1.3, 2)
total = round((video_cost + fixed_cost) * profit_multiplier, 2)
若传入 actual_tokens 则用它替代计算值。
"""
prefix = _match_model_prefix(model)
res_key = _infer_resolution_key(int(height or 720))
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)
if actual_tokens is not None and actual_tokens > 0:
tokens = float(actual_tokens)
else:
dur = max(1, int(duration_seconds or 15))
w = max(1, int(width or 1))
h = max(1, int(height or 1))
tokens = dur * w * h * int(fps or VIRAL_VIDEO_FPS) / 1024.0
video_cost = tokens / 1_000_000.0 * float(price)
total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER
return round(float(total), 2)
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(显示名称)
# 说明:仅保留需要扣点的场景;免费场景不要写入此字典。
# 每个场景: base_points(基础积分), unit(计费单位), name(显示名称), dynamic(是否动态定价)
# 说明:爆款视频(viral_video)走动态定价(预扣→结算多退少补),因此不使用 @points_gate
# 装饰器,base_points=0,dynamic=True;前端展示场景列表时仍可看到。
POINTS_SCENES: dict[str, dict] = {
"voice_clone_train": {
@@ -140,6 +191,13 @@ POINTS_SCENES: dict[str, dict] = {
"name": "声音克隆合成",
"description": "克隆音色合成每分钟消耗 1 积分",
},
"viral_video": {
"base_points": 0,
"unit": "次",
"name": "爆款视频",
"dynamic": True,
"description": "爆款视频动态定价(按视频时长/分辨率/模型计算,预扣→结算多退少补)",
},
}
# 免费用户积分消耗上浮系数(仅对 voice_clone_synth 生效)
@@ -179,20 +237,26 @@ def calculate_points_cost(
"""计算指定场景的积分消耗。
Args:
scene_key: 场景标识(当前仅支持 voice_clone_train/voice_clone_synth)
scene_key: 场景标识(当前支持 voice_clone_train/voice_clone_synth/viral_video;
viral_video 为动态定价场景,此处返回 0,由业务侧调用
calculate_viral_video_credits 手动计算)
is_member: 是否付费会员
quantity: 数量(按次计费场景)
duration_minutes: 时长分钟数(按时长计费场景)
member_type: 会员类型 (monthly/quarterly/yearly),用于折扣
Returns:
实际消耗积分(float;已含免费用户 ×1.15 上浮或会员折扣);免费/已下线场景统一返回 0。
实际消耗积分(float;已含免费用户 ×1.15 上浮或会员折扣);免费/动态/已下线场景统一返回 0。
"""
scene = POINTS_SCENES.get(scene_key)
if not scene:
# 已下线/未注册的场景统一返回 0(免费),保持向后兼容
return 0.0
# 动态定价场景(如 viral_video)由业务侧手动计算,这里统一返回 0
if scene.get("dynamic"):
return 0.0
base = scene["base_points"]
if base == 0:
return 0.0
+3 -1
View File
@@ -129,7 +129,9 @@ class TestPointsRulesDescription:
from app.api.routes.points import get_rules
resp = get_rules(_current_user=_make_cu())
assert len(resp.rules) == 2
# 场景列表包含 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)
for rule in resp.rules:
assert rule.description, f"{rule.scene_key} missing description"
assert isinstance(rule.description, str)
+76 -2
View File
@@ -19,9 +19,22 @@ 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 set(POINTS_SCENES.keys()) == {"voice_clone_train", "voice_clone_synth"}
# 保留声音克隆两个场景
assert "voice_clone_train" in POINTS_SCENES
assert "voice_clone_synth" in POINTS_SCENES
def test_required_keys_present(self):
for key, scene in POINTS_SCENES.items():
@@ -36,6 +49,11 @@ class TestPointsScenesConfig:
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
class TestPointsPackages:
def test_three_packages(self):
@@ -392,3 +410,59 @@ class TestCalculateViralVideoCredits:
c_none = calculate_viral_video_credits(10, 1280, 720, fps=None)
assert c_zero == c_default
assert c_none == c_default
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
+226 -9
View File
@@ -14,6 +14,8 @@ 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))
@@ -136,6 +138,163 @@ 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 ──────────────────────────────────────────────────────
@@ -591,44 +750,62 @@ class TestConfirmCopyPointsDeduction:
class TestEstimateCredits:
"""POST /estimate-credits: 纯计算预估积分。"""
def test_estimate_returns_float(self):
"""正常参数应返回 estimated_credits 为 float 且>0。"""
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)
# 不需要 db / user 之外的依赖;authenticated_user 仍要传
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(self):
"""estimate_credits 应调用 resolve_video_dimensions 和 calculate_viral_video_credits。"""
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", return_value=8.88) as mock_calc,
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()
# 传给 calculate 的参数应包含 duration=20, w=1920, h=1080, model="seedance-2.5"
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'。"""
@@ -637,13 +814,53 @@ class TestEstimateCredits:
req = EstimateCreditsRequest(model="", resolution="720p", ratio="9:16", duration=10)
user = _auth_user("u1")
fake_bd = {
"tokens": 500.0, "video_cost": 0.5, "fixed_cost": 0.15,
"profit_multiplier": 1.3, "model_price": 70.0,
"width": 720, "height": 1280, "fps": 24,
}
with (
patch("packages.domain.points_rules.resolve_video_dimensions", return_value=(720, 1280)),
patch("packages.domain.points_rules.calculate_viral_video_credits", return_value=3.5) as mock_calc,
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"