From 5d6a4675fb8108941b29cd88c37752db045f9dac Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sat, 3 Oct 2026 00:20:20 +0800 Subject: [PATCH] =?UTF-8?q?feat(viral-video):=20=E5=8A=A8=E6=80=81?= =?UTF-8?q?=E7=A7=AF=E5=88=86=E5=AE=9A=E4=BB=B7=EF=BC=88=E6=8C=89tokens?= =?UTF-8?q?=C3=97=E5=8D=95=E4=BB=B7=C3=971.3=EF=BC=8C=E4=BF=9D=E7=95=99?= =?UTF-8?q?=E4=B8=A4=E4=BD=8D=E5=B0=8F=E6=95=B0=EF=BC=89=20(#2152)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-authored-by: xiaoxia Co-committed-by: xiaoxia --- .../093_viral_video_pricing_points_float.py | 87 ++++++ apps/api/app/api/routes/viral_video.py | 59 +++- apps/api/app/schemas/points.py | 20 +- apps/api/app/schemas/viral_video.py | 26 +- apps/worker/worker_app/tasks/viral_video.py | 139 ++++++++- packages/adapters/sqlalchemy_impl/models.py | 17 +- .../sqlalchemy_impl/viral_video_repository.py | 15 +- packages/domain/points_account.py | 6 +- packages/domain/points_rules.py | 126 +++++++- packages/domain/points_service.py | 100 ++++++- packages/domain/viral_video.py | 7 +- packages/shared/ai_client.py | 13 +- packages/shared/ai_service.py | 6 +- tests/unit/test_2035_coverage.py | 4 +- tests/unit/test_ai_client_video.py | 20 +- tests/unit/test_points_rules.py | 275 ++++++++++++++++++ tests/unit/test_points_service.py | 169 +++++++++++ tests/unit/test_viral_video.py | 8 +- tests/unit/test_viral_video_p0.py | 16 +- tests/unit/test_viral_video_routes.py | 250 +++++++++++++++- 20 files changed, 1283 insertions(+), 80 deletions(-) create mode 100644 alembic/versions/093_viral_video_pricing_points_float.py diff --git a/alembic/versions/093_viral_video_pricing_points_float.py b/alembic/versions/093_viral_video_pricing_points_float.py new file mode 100644 index 000000000..4b3d54bb5 --- /dev/null +++ b/alembic/versions/093_viral_video_pricing_points_float.py @@ -0,0 +1,87 @@ +"""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 diff --git a/apps/api/app/api/routes/viral_video.py b/apps/api/app/api/routes/viral_video.py index 564ecada1..b55e06bbe 100644 --- a/apps/api/app/api/routes/viral_video.py +++ b/apps/api/app/api/routes/viral_video.py @@ -32,6 +32,8 @@ from app.schemas.viral_video import ( ConfirmCopyRequest, ConfirmIntentRequest, CreateViralVideoRequest, + EstimateCreditsRequest, + EstimateCreditsResponse, GenerateCopyRequest, StyleTemplateListResponse, StyleTemplateResponse, @@ -140,7 +142,9 @@ def _to_response(job) -> ViralVideoJobResponse: video_model=getattr(job, "video_model", "") or "", intent_result=job.intent_result, result_video_url=job.result_video_url, - credits_cost=job.credits_cost, + 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), error_msg=job.error_msg, retry_count=job.retry_count, started_at=job.started_at, @@ -193,6 +197,7 @@ 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, ) @@ -235,6 +240,7 @@ 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) @@ -298,6 +304,7 @@ 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) @@ -330,6 +337,41 @@ 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) @@ -344,6 +386,21 @@ 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: + """爆款视频积分预估(纯计算,不扣费、不创建任务)。""" + 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" + ) + return EstimateCreditsResponse(estimated_credits=credits) + + @router.get("/history", response_model=ViralVideoHistoryResponse) def list_viral_video_history( limit: int = 50, diff --git a/apps/api/app/schemas/points.py b/apps/api/app/schemas/points.py index 6424b28a1..57ccd25e2 100644 --- a/apps/api/app/schemas/points.py +++ b/apps/api/app/schemas/points.py @@ -13,9 +13,9 @@ from pydantic import BaseModel, Field class PointsBalanceResponse(BaseModel): """积分余额 + 会员状态""" - balance: int = Field(..., description="当前积分余额") - total_earned: int = Field(..., description="累计获得积分") - total_spent: int = Field(..., description="累计消耗积分") + balance: float = Field(..., description="当前积分余额") + total_earned: float = Field(..., description="累计获得积分") + total_spent: float = 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: int - balance_after: int + amount: float + balance_after: float description: str = "" ref_id: str = "" created_at: Optional[str] = None @@ -99,9 +99,9 @@ class PointsCheckResponse(BaseModel): """消费前余额检查响应""" allowed: bool - required_points: int - current_balance: int - remaining_after: int + required_points: float + current_balance: float + remaining_after: float is_free_quota: bool = False @@ -112,7 +112,7 @@ class PointsDeductRequest(BaseModel): """积分扣减请求""" scene_key: str - amount: int + amount: float 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: int + points_balance: float max_resolution: str = Field( default="1080p", description="可用最高分辨率: 720p(free) / 1080p(paid)", diff --git a/apps/api/app/schemas/viral_video.py b/apps/api/app/schemas/viral_video.py index 96bd65843..06d4d0375 100755 --- a/apps/api/app/schemas/viral_video.py +++ b/apps/api/app/schemas/viral_video.py @@ -22,6 +22,7 @@ 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) -- @@ -88,6 +89,7 @@ class CreateViralVideoRequest(BaseModel): voice_source: str = "" video_ratio: str = "9:16" video_model: str = "" + video_resolution: str = "720p" @field_validator("fusion_level") @classmethod @@ -117,6 +119,7 @@ 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) @@ -141,6 +144,7 @@ class GenerateCopyRequest(BaseModel): voice_source: str = "" video_ratio: str = "9:16" video_model: str = "" + video_resolution: str = "720p" @field_validator("fusion_level") @classmethod @@ -218,7 +222,9 @@ class ViralVideoJobResponse(BaseModel): video_model: str = "" intent_result: dict | None = None result_video_url: str = "" - credits_cost: int = 0 + video_resolution: str = "720p" + credits_prepaid: float = 0.0 + credits_cost: float = 0.0 error_msg: str = "" retry_count: int = 0 started_at: datetime | None = None @@ -250,6 +256,24 @@ class AnalyzeStyleResponse(BaseModel): style_guide: dict | None = None +# -- 积分预估 -- + + +class EstimateCreditsRequest(BaseModel): + """爆款视频积分预估请求。""" + + model: str = "" + resolution: str = "720p" + ratio: str = "9:16" + duration: int = Field(default=15, ge=5, le=30) + + +class EstimateCreditsResponse(BaseModel): + """爆款视频积分预估响应。""" + + estimated_credits: float + + # -- WebSocket 事件 Schema -- diff --git a/apps/worker/worker_app/tasks/viral_video.py b/apps/worker/worker_app/tasks/viral_video.py index e5f50315f..8adc448c2 100644 --- a/apps/worker/worker_app/tasks/viral_video.py +++ b/apps/worker/worker_app/tasks/viral_video.py @@ -37,7 +37,6 @@ from packages.adapters.sqlalchemy_impl.viral_video_repository import ( SQLAlchemyViralVideoJobRepository, ) from packages.domain.viral_video import ( - CREDITS_VIRAL_VIDEO_COST, STAGE_LABELS, ViralVideoJob, ViralVideoStage, @@ -1110,14 +1109,18 @@ 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) -> str: - """步骤 6: v1.6 单次 Seedance 生成(不再分段/拼接)。""" +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.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" # reference_audios: TTS 音频驱动口型 ref_audios = [tts_audio_url] if tts_audio_url else [] @@ -1141,12 +1144,12 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non ) logger.info("[爆款视频] Seedance prompt (前300字): %s", prompt[:300]) - video_path = call_video_generation( + result = call_video_generation( prompt=prompt, image_url=first_image, duration=dur, ratio=ratio, - resolution="720p", + resolution=resolution, output_dir=str(tmpdir), model=model, generate_audio=True, # Seedance 原生生成环境音效/BGM;口型由 reference_audios 的 TTS 驱动 @@ -1154,10 +1157,16 @@ def _step_render(job: ViralVideoJob, copy_result: dict, tts_audio_url: str | Non reference_audios=ref_audios, reference_videos=ref_videos, ) + if not result or not isinstance(result, dict): + raise RuntimeError("Seedance 视频生成失败:返回为空") + 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("Seedance 视频生成失败:返回空文件或路径不存在") - logger.info("[爆款视频] Seedance 单次生成完成: %s size=%d", video_path, Path(video_path).stat().st_size) - return str(video_path) + logger.info( + "[爆款视频] Seedance 单次生成完成: %s size=%d usage=%s", video_path, Path(video_path).stat().st_size, usage + ) + return str(video_path), (usage if isinstance(usage, dict) else None) def _step_upload(job: ViralVideoJob, video_path: str) -> str: @@ -1555,6 +1564,90 @@ 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", + ) + actual_tokens = None + if isinstance(usage, dict): + at = usage.get("completion_tokens") + if isinstance(at, (int, float)) and at > 0: + actual_tokens = int(at) + actual_credits = calculate_viral_video_credits( + int(getattr(job, "duration", 15) or 15), + 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。 @@ -1596,9 +1689,17 @@ 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 = _step_render(job, copy_result, tts_url) + 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 _emit_progress(job_id, ViralVideoStage.RENDERING, 92.0, "视频生成完成") # Step 7: Upload @@ -1610,7 +1711,8 @@ def _run_render_pipeline(job_id: str, session, repo, job) -> dict: if video_url: _wait_oss_ready(video_url, timeout_sec=10) - job.credits_cost = CREDITS_VIRAL_VIDEO_COST + # 积分结算:按实际 tokens 多退少补 + _settle_viral_video(job, usage) job.mark_completed(video_url) job.current_stage = ViralVideoStage.UPLOADING job.phase_message = "视频生成完成" @@ -1648,6 +1750,23 @@ 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: diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 6fc5609a3..993c21c49 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -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(Integer, nullable=False, default=0) + points_balance = Column(Float, 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(Integer, nullable=False, default=0) - total_earned = Column(Integer, nullable=False, default=0) - total_spent = Column(Integer, nullable=False, default=0) + balance = Column(Float, nullable=False, default=0) + total_earned = Column(Float, nullable=False, default=0) + total_spent = Column(Float, 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(Integer, nullable=False) - balance_after = Column(Integer, nullable=False) + amount = Column(Float, nullable=False) + balance_after = Column(Float, 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,7 +961,10 @@ 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(Integer, nullable=False, default=0) + 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="") error_msg = Column(Text, nullable=False, default="") retry_count = Column(Integer, nullable=False, default=0) started_at = Column(DateTime(timezone=True), nullable=True) diff --git a/packages/adapters/sqlalchemy_impl/viral_video_repository.py b/packages/adapters/sqlalchemy_impl/viral_video_repository.py index df96ad0f7..e137b037b 100755 --- a/packages/adapters/sqlalchemy_impl/viral_video_repository.py +++ b/packages/adapters/sqlalchemy_impl/viral_video_repository.py @@ -46,7 +46,10 @@ 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 "", - credits_cost=model.credits_cost or 0, + 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), error_msg=model.error_msg or "", retry_count=model.retry_count or 0, started_at=model.started_at, @@ -95,7 +98,10 @@ class SQLAlchemyViralVideoJobRepository: generated_copy_text=job.generated_copy_text, copy_result=job.copy_result, result_video_url=job.result_video_url, - credits_cost=job.credits_cost, + 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), error_msg=job.error_msg, retry_count=job.retry_count, started_at=job.started_at, @@ -121,7 +127,10 @@ 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.credits_cost = job.credits_cost + 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.error_msg = job.error_msg model.retry_count = job.retry_count model.started_at = job.started_at diff --git a/packages/domain/points_account.py b/packages/domain/points_account.py index 3654aeb79..476ce58e6 100644 --- a/packages/domain/points_account.py +++ b/packages/domain/points_account.py @@ -9,9 +9,9 @@ from uuid import uuid4 class PointsAccount: id: str user_id: str - balance: int = 0 - total_earned: int = 0 - total_spent: int = 0 + balance: float = 0.0 + total_earned: float = 0.0 + total_spent: float = 0.0 created_at: datetime = field(default_factory=lambda: datetime.now(UTC)) updated_at: datetime = field(default_factory=lambda: datetime.now(UTC)) diff --git a/packages/domain/points_rules.py b/packages/domain/points_rules.py index 13b3df8f5..f62fda66e 100644 --- a/packages/domain/points_rules.py +++ b/packages/domain/points_rules.py @@ -2,13 +2,127 @@ v1.6.1: 按产品决策,智能混剪/AI数字人/AI配音/抖音解析/改写/标题/封面 全部免费, 仅保留声音克隆合成(voice_clone_synth)的扣点逻辑;声音克隆训练保持免费。 -爆款视频(viral_video)后续走动态定价,暂不加入本文件。 +爆款视频(viral_video)走动态定价,见本文件 VIRAL_VIDEO_MODEL_PRICES + calculate_viral_video_credits。 """ from __future__ import annotations import math +# ============ 爆款视频动态定价 (#2151) ============ +# key = (model_id, resolution, has_video_input),单位:元/百万token +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, +} + +# 固定成本(元):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", +} +# 分辨率 -> 高度 +_RESOLUTION_HEIGHT: dict[str, int] = {"480p": 480, "720p": 720, "1080p": 1080} + + +def resolve_video_dimensions(resolution: str, ratio: str) -> tuple[int, int]: + """把 (resolution, ratio) 解析为 (width, height)。""" + key = str(resolution or "").strip() + key_l = key.lower() + res_key = _RESOLUTION_ALIASES.get(key_l) or _RESOLUTION_ALIASES.get(key) or "720p" + h = _RESOLUTION_HEIGHT.get(res_key, 720) + r = str(ratio or "").strip().lower() + if r == "16:9": + w = h * 16 // 9 + elif r == "1:1": + w = h + else: + w = h * 9 // 16 + return int(w), int(h) + + +def _match_model_prefix(model: str) -> str: + """匹配 model 前缀。""" + m = (model or "").strip().lower() + for prefix in ("seedance-2.5", "seedance-2.0"): + if m.startswith(prefix): + return prefix + return "seedance-2.5" + + +def _infer_resolution_key(height: int) -> str: + """从像素高度推断 resolution key。""" + if height >= 1000: + return "1080p" + if height >= 650: + return "720p" + return "480p" + + +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 元)。 + + 公式: + tokens = duration * width * height * fps / 1024 + video_cost = tokens / 1_000_000 * model_token_price + total = round((video_cost + fixed_cost) * 1.3, 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) + + # ============ 场景定义 ============ # 每个场景: base_points(基础积分), unit(计费单位), name(显示名称) # 说明:仅保留需要扣点的场景;免费场景不要写入此字典。 @@ -61,7 +175,7 @@ def calculate_points_cost( quantity: int = 1, duration_minutes: float = 0, member_type: str | None = None, -) -> int: +) -> float: """计算指定场景的积分消耗。 Args: @@ -72,16 +186,16 @@ def calculate_points_cost( member_type: 会员类型 (monthly/quarterly/yearly),用于折扣 Returns: - 实际消耗积分(已含免费用户 ×1.15 上浮或会员折扣);免费/已下线场景统一返回 0。 + 实际消耗积分(float;已含免费用户 ×1.15 上浮或会员折扣);免费/已下线场景统一返回 0。 """ scene = POINTS_SCENES.get(scene_key) if not scene: # 已下线/未注册的场景统一返回 0(免费),保持向后兼容 - return 0 + return 0.0 base = scene["base_points"] if base == 0: - return 0 + return 0.0 unit = scene["unit"] if unit == "分钟": @@ -96,4 +210,4 @@ def calculate_points_cost( elif not is_member: total_base = math.ceil(total_base * FREE_USER_MULTIPLIER) - return total_base + return float(total_base) diff --git a/packages/domain/points_service.py b/packages/domain/points_service.py index 4259f4f13..935df002d 100644 --- a/packages/domain/points_service.py +++ b/packages/domain/points_service.py @@ -83,7 +83,7 @@ class PointsService: # ──────────────── 余额检查 ──────────────── - def check_balance(self, user_id: str, amount: int, db: Session) -> dict[str, Any]: + def check_balance(self, user_id: str, amount: float, db: Session) -> dict[str, Any]: """检查余额是否足够。""" account_data = self.get_or_create_account(user_id, db) balance = account_data["balance"] @@ -99,7 +99,7 @@ class PointsService: def deduct_points( self, user_id: str, - amount: int, + amount: float, source: str, db: Session, description: str = "", @@ -108,7 +108,7 @@ class PointsService: """扣减积分(事务性:SELECT FOR UPDATE → 检查余额 → 扣减 → 流水 → 同步用户表)。 Returns: - {"success": True/False, "balance": int, "transaction_id": str|None} + {"success": True/False, "balance": float, "transaction_id": str|None} """ PointsAccountModel, PointsTransactionModel, _, _, UserModel = _get_models() @@ -172,7 +172,7 @@ class PointsService: except Exception: db.rollback() logger.exception( - "积分扣减失败: user_id=%s, amount=%d, source=%s", + "积分扣减失败: user_id=%s, amount=%.2f, source=%s", user_id, amount, source, @@ -184,7 +184,7 @@ class PointsService: def add_points( self, user_id: str, - amount: int, + amount: float, source: str, db: Session, description: str = "", @@ -241,7 +241,7 @@ class PointsService: except Exception: db.rollback() logger.exception( - "积分增加失败: user_id=%s, amount=%d, source=%s", + "积分增加失败: user_id=%s, amount=%.2f, source=%s", user_id, amount, source, @@ -253,7 +253,7 @@ class PointsService: def refund_points( self, user_id: str, - amount: int, + amount: float, source: str, db: Session, ref_id: str = "", @@ -269,6 +269,92 @@ 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( diff --git a/packages/domain/viral_video.py b/packages/domain/viral_video.py index 5155c306b..73f1665f4 100755 --- a/packages/domain/viral_video.py +++ b/packages/domain/viral_video.py @@ -69,8 +69,6 @@ class PromptType(StrEnum): STYLE_CONSTRAINT = "style_constraint" -CREDITS_VIRAL_VIDEO_COST = 50 - STAGE_LABELS = { ViralVideoStage.IMAGE_ANALYSIS: "图片分析", ViralVideoStage.VIDEO_ANALYSIS: "视频风格分析", @@ -121,7 +119,10 @@ class ViralVideoJob: phase_message: str = "" # 阶段中文提示文案,前端轮询直接展示 heartbeat_at: datetime | None = None # worker 心跳时间,用于超时僵尸任务检测 result_video_url: str = "" - credits_cost: int = 0 + video_resolution: str = "720p" + credits_prepaid: float = 0.0 + credits_transaction_id: str = "" + credits_cost: float = 0.0 error_msg: str = "" retry_count: int = 0 started_at: datetime | None = None diff --git a/packages/shared/ai_client.py b/packages/shared/ai_client.py index 819aa61ba..145be55f0 100755 --- a/packages/shared/ai_client.py +++ b/packages/shared/ai_client.py @@ -261,8 +261,11 @@ class DoubaoClient: reference_images: list[str] | None = None, reference_audios: list[str] | None = None, reference_videos: list[str] | None = None, - ) -> str | None: - """调用 Seedance 2.5 生视频(异步任务→轮询→下载),返回本地 MP4 路径;失败返回 None。 + ) -> dict | None: + """调用 Seedance 2.5 生视频(异步任务→轮询→下载)。 + + 成功返回 {"video_path": str, "usage": dict | None},失败返回 None。 + usage 是 Seedance 返回的计费信息(含 completion_tokens)。 【v1.6.1 修复】严格按官方 content 数组协议构造请求: - 所有参考(图/音/视)必须放进 content 数组并带 role 字段,不能放顶层 reference_audios/reference_videos(非官方字段,会被忽略或导致异常)。 @@ -420,6 +423,7 @@ 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 while time.time() < deadline: @@ -437,8 +441,9 @@ class DoubaoClient: 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", task_id, poll_count) + logger.info("Seedance 任务成功: task_id=%s polls=%d usage=%s", task_id, poll_count, usage) break last_err = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}") logger.error("Seedance succeeded 但无 video_url: %s", last_err) @@ -500,7 +505,7 @@ class DoubaoClient: except Exception: pass return None - return local_path + return {"video_path": local_path, "usage": usage} except Exception as e: logger.error("Seedance 视频下载失败: %s", e, exc_info=True) return None diff --git a/packages/shared/ai_service.py b/packages/shared/ai_service.py index 4c0aa7791..9881f175b 100755 --- a/packages/shared/ai_service.py +++ b/packages/shared/ai_service.py @@ -617,8 +617,10 @@ def call_video_generation( reference_images: list[str] | None = None, reference_audios: list[str] | None = None, reference_videos: list[str] | None = None, -) -> str | None: - """调用 Seedance 2.5 生成视频(v1.6.1 单次出片版),返回本地 MP4 路径;失败返回 None。 +) -> dict | None: + """调用 Seedance 2.5 生成视频(v1.6.1 单次出片版)。 + + 成功返回 {"video_path": str, "usage": dict | None}(usage 含 completion_tokens),失败返回 None。 v1.6.1 关键约束(避免 20min 卡死): - 参考音频/视频/多图全部放进 content 数组并带 role=reference_audio/reference_video/reference_image; diff --git a/tests/unit/test_2035_coverage.py b/tests/unit/test_2035_coverage.py index b8bbedf6d..7c682b876 100644 --- a/tests/unit/test_2035_coverage.py +++ b/tests/unit/test_2035_coverage.py @@ -87,7 +87,9 @@ _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:文案关键词") - end = src.index("from packages.middleware") + # 用紧跟 _infer_expected_categories 后的 logger 行作为结束锚点 + end_marker = "\nlogger = logging.getLogger" + end = src.index(end_marker, start) code = src[start:end] ns: dict = {} exec(code, ns) diff --git a/tests/unit/test_ai_client_video.py b/tests/unit/test_ai_client_video.py index dfdb188bb..1cf01e336 100644 --- a/tests/unit/test_ai_client_video.py +++ b/tests/unit/test_ai_client_video.py @@ -112,10 +112,10 @@ class TestVideoGenerationHappyPath: resolution="720p", output_dir=str(tmp_path), ) - 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 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 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 - assert Path(out).read_bytes() == b"DATA" + assert out is not None and isinstance(out, dict) + assert Path(out["video_path"]).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 - assert Path(out).exists() + assert out is not None and isinstance(out, dict) + assert Path(out["video_path"]).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 - assert "/tmp/seedance_t-default_00000001.mp4" in out + assert out is not None and isinstance(out, dict) + assert out["video_path"] == "/tmp/seedance_t-default_00000001.mp4" assert captured["json"]["generate_audio"] is True assert captured["json"]["watermark"] is True assert captured["json"]["ratio"] == "1:1" diff --git a/tests/unit/test_points_rules.py b/tests/unit/test_points_rules.py index 17294bfc1..6dc7682f3 100644 --- a/tests/unit/test_points_rules.py +++ b/tests/unit/test_points_rules.py @@ -117,3 +117,278 @@ class TestCalculatePointsCost: 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 → h=480, w 按 16//9 计算。""" + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions("480p", "16:9") + assert h == 480 + assert w == 480 * 16 // 9 + + 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("2160p", "1:1") + assert h == 720 + assert w == 720 + + def test_empty_resolution_defaults_to_720p_9_16(self): + """空 resolution + 空 ratio → 默认 720p + 9:16。""" + from packages.domain.points_rules import resolve_video_dimensions + + w, h = resolve_video_dimensions("", "") + assert h == 720 + assert w == 720 * 9 // 16 + + 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 h == 720 + assert w == 720 * 9 // 16 + + 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(): 1000+/650-999/<650 三个分支。""" + + def test_height_ge_1000_is_1080p(self): + from packages.domain.points_rules import _infer_resolution_key + + assert _infer_resolution_key(1000) == "1080p" + assert _infer_resolution_key(1080) == "1080p" + assert _infer_resolution_key(2160) == "1080p" + + def test_height_650_to_999_is_720p(self): + from packages.domain.points_rules import _infer_resolution_key + + assert _infer_resolution_key(650) == "720p" + assert _infer_resolution_key(720) == "720p" + assert _infer_resolution_key(999) == "720p" + + def test_height_lt_650_is_480p(self): + from packages.domain.points_rules import _infer_resolution_key + + assert _infer_resolution_key(480) == "480p" + assert _infer_resolution_key(649) == "480p" + assert _infer_resolution_key(0) == "480p" + + +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 diff --git a/tests/unit/test_points_service.py b/tests/unit/test_points_service.py index 3d744801d..4bfe239c4 100644 --- a/tests/unit/test_points_service.py +++ b/tests/unit/test_points_service.py @@ -179,3 +179,172 @@ 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): + """actualestimated 且补扣成功 → 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 diff --git a/tests/unit/test_viral_video.py b/tests/unit/test_viral_video.py index dcb1f0ba8..a7d363cf7 100755 --- a/tests/unit/test_viral_video.py +++ b/tests/unit/test_viral_video.py @@ -17,7 +17,6 @@ import pytest from pydantic import ValidationError from packages.domain.viral_video import ( - CREDITS_VIRAL_VIDEO_COST, STAGE_LABELS, FusionLevel, StyleStrength, @@ -129,9 +128,6 @@ 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: """阶段枚举测试。""" @@ -559,7 +555,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" + mock_render.return_value = ("/tmp/video.mp4", {"completion_tokens": 1000000}) mock_upload.return_value = "https://oss.example.com/final.mp4" result = resume_viral_video_pipeline.run("job-001") @@ -567,4 +563,4 @@ class TestPipelineIntegration: assert result["ok"] is True assert result["video_url"] == "https://oss.example.com/final.mp4" assert job.status == ViralVideoStatus.COMPLETED - assert job.credits_cost == CREDITS_VIRAL_VIDEO_COST + assert isinstance(job.credits_cost, float) diff --git a/tests/unit/test_viral_video_p0.py b/tests/unit/test_viral_video_p0.py index acb8233cf..2ab4de7c7 100644 --- a/tests/unit/test_viral_video_p0.py +++ b/tests/unit/test_viral_video_p0.py @@ -212,10 +212,13 @@ 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 = str(out) + mock_client.video_generation.return_value = { + "video_path": str(out), + "usage": {"completion_tokens": 1000000}, + } mock_get.return_value = mock_client result = call_video_generation(prompt="测试", image_url="https://img/x.jpg", duration=5, ratio="9:16") - assert result == str(out) + assert result is not None and result["video_path"] == str(out) mock_client.video_generation.assert_called_once() kwargs = mock_client.video_generation.call_args.kwargs assert kwargs["prompt"] == "测试" @@ -238,7 +241,10 @@ 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 = str(out) + mock_client.video_generation.return_value = { + "video_path": str(out), + "usage": {"completion_tokens": 1500000}, + } mock_get.return_value = mock_client result = call_video_generation( prompt="测试", @@ -251,7 +257,7 @@ class TestCallVideoGenerationV16: generate_audio=True, model="doubao-seedance-2-5-260628", ) - assert result == str(out) + assert result is not None and result["video_path"] == 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}" @@ -270,7 +276,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 = str(out) + mock_client.video_generation.return_value = {"video_path": str(out), "usage": None} mock_get.return_value = mock_client call_video_generation(prompt="测试", duration=10, ratio="16:9") kwargs = mock_client.video_generation.call_args.kwargs diff --git a/tests/unit/test_viral_video_routes.py b/tests/unit/test_viral_video_routes.py index 3dbbdcfdf..d0d9a9629 100644 --- a/tests/unit/test_viral_video_routes.py +++ b/tests/unit/test_viral_video_routes.py @@ -60,7 +60,10 @@ def _make_job(job_id: str = "job-1", user_id: str = "u1", status: str = "pending "voice_mode": "global", "video_ratio": "9:16", "video_model": "", - "credits_cost": 0, + "video_resolution": "720p", + "credits_prepaid": 0.0, + "credits_transaction_id": "", + "credits_cost": 0.0, "current_stage": "", "phase_message": "", "updated_at": None, @@ -399,3 +402,248 @@ 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(self): + """正常参数应返回 estimated_credits 为 float 且>0。""" + 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 + + def test_estimate_uses_dimensions_resolver(self): + """estimate_credits 应调用 resolve_video_dimensions 和 calculate_viral_video_credits。""" + 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") + + 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, + ): + 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 + + 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") + + 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, + ): + 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