diff --git a/apps/api/app/api/routes/viral_video.py b/apps/api/app/api/routes/viral_video.py index b55e06bbe..fa37af60b 100644 --- a/apps/api/app/api/routes/viral_video.py +++ b/apps/api/app/api/routes/viral_video.py @@ -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}") diff --git a/apps/api/app/schemas/viral_video.py b/apps/api/app/schemas/viral_video.py index 06d4d0375..cdc8d758a 100755 --- a/apps/api/app/schemas/viral_video.py +++ b/apps/api/app/schemas/viral_video.py @@ -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 -- diff --git a/packages/domain/points_rules.py b/packages/domain/points_rules.py index f62fda66e..d1a53697b 100644 --- a/packages/domain/points_rules.py +++ b/packages/domain/points_rules.py @@ -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 diff --git a/tests/unit/test_points_routes.py b/tests/unit/test_points_routes.py index 83f3e5ed3..2b8a7f08b 100644 --- a/tests/unit/test_points_routes.py +++ b/tests/unit/test_points_routes.py @@ -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) diff --git a/tests/unit/test_points_rules.py b/tests/unit/test_points_rules.py index 6dc7682f3..d3b29598a 100644 --- a/tests/unit/test_points_rules.py +++ b/tests/unit/test_points_rules.py @@ -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 diff --git a/tests/unit/test_viral_video_routes.py b/tests/unit/test_viral_video_routes.py index d0d9a9629..d81c5e841 100644 --- a/tests/unit/test_viral_video_routes.py +++ b/tests/unit/test_viral_video_routes.py @@ -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"