feat(viral-video): 动态积分定价(按tokens×单价×1.3,保留两位小数) (#2152)
CI/CD Pipeline / Check if frontend-only change (push) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (push) Successful in 1s
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 3s
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 4s
CI/CD Pipeline / Check push changed paths (push) Successful in 15s
CI/CD Pipeline / Frontend Lint (push) Has been skipped
CI/CD Pipeline / PR Build API Image (push) Has been skipped
CI/CD Pipeline / PR Build Web Image (push) Has been skipped
CI/CD Pipeline / PR Build Worker Image (push) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Validate - Style (pull_request) Has been skipped
CI/CD Pipeline / Validate - Security (pull_request) Has been skipped
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Has been skipped
CI/CD Pipeline / Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 52s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 54s
CI/CD Pipeline / Build Staging Worker Image (push) Successful in 31s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 2s
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m49s
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (push) Successful in 55s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m17s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (push) Successful in 3m23s
CI/CD Pipeline / Retag skipped Staging API Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (push) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 4m19s
CI/CD Pipeline / Frontend Unit Tests (push) Successful in 4m30s
CI/CD Pipeline / Validate - Python (mypy + alembic) (push) Successful in 4m47s
CI/CD Pipeline / Validate - Style (push) Successful in 5m10s
CI/CD Pipeline / Validate - Security (push) Has been cancelled
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production API Image (push) Has been cancelled
CI/CD Pipeline / Build Production Web Image (push) Has been cancelled
CI/CD Pipeline / Build Production Worker Image (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
CI/CD Pipeline / ACR Image Cleanup (push) Has been cancelled
CI/CD Pipeline / Canary Release to Production (push) Has been cancelled
CI/CD Pipeline / CI Gate (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
AI Code Review / AI Code Review (pull_request) Has been cancelled

Co-authored-by: xiaoxia <dev@xiaoxiajianji.com>
Co-committed-by: xiaoxia <dev@xiaoxiajianji.com>
This commit was merged in pull request #2152.
This commit is contained in:
2026-10-03 00:20:20 +08:00
committed by auto-approve-bot
parent 774845bf91
commit 5d6a4675fb
20 changed files with 1283 additions and 80 deletions
@@ -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
+58 -1
View File
@@ -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,
+10 -10
View File
@@ -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)",
+25 -1
View File
@@ -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 --
+129 -10
View File
@@ -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:
+10 -7
View File
@@ -56,7 +56,7 @@ class UserModel(Base):
is_member = Column(Boolean, nullable=False, default=False)
member_type = Column(String(20), nullable=True)
member_expires_at = Column(DateTime, nullable=True)
points_balance = Column(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)
@@ -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
+3 -3
View File
@@ -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))
+120 -6
View File
@@ -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)
+93 -7
View File
@@ -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(
+4 -3
View File
@@ -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
+9 -4
View File
@@ -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
+4 -2
View File
@@ -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;
+3 -1
View File
@@ -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)
+10 -10
View File
@@ -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"
+275
View File
@@ -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
+169
View File
@@ -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):
"""actual<estimated 时走 refund_points,返回 action=refund。"""
refund_res = {"success": True, "balance": 60.0, "transaction_id": "tr-1"}
with patch.object(service, "refund_points", return_value=refund_res) as mock_refund:
result = service.settle_viral_video(user_id, estimated=20.0, actual=15.0, txn_id="t2", db=db_session)
assert result["success"] is True
assert result["action"] == "refund"
assert result["amount"] == 5.0
assert result["diff"] == -5.0
mock_refund.assert_called_once()
rk = mock_refund.call_args.kwargs
assert rk["user_id"] == user_id
assert rk["amount"] == 5.0
assert rk["source"] == "viral_video"
assert rk["ref_id"] == "t2"
assert rk["description"] == "爆款视频结算退费"
def test_refund_exception_returns_failure(self, service, db_session, user_id):
"""refund_points 抛异常时,应捕获并返回 success=False。"""
with patch.object(service, "refund_points", side_effect=RuntimeError("db down")):
result = service.settle_viral_video(user_id, estimated=20.0, actual=10.0, txn_id="t3", db=db_session)
assert result["success"] is False
assert result["action"] == "refund"
def test_deduct_when_actual_greater_than_estimated_success(self, service, db_session, user_id):
"""actual>estimated 且补扣成功 → action=deduct, success=True。"""
deduct_res = {"success": True, "balance": 40.0, "transaction_id": "td-1"}
with patch.object(service, "deduct_points", return_value=deduct_res) as mock_deduct:
result = service.settle_viral_video(user_id, estimated=15.0, actual=20.0, txn_id="t4", db=db_session)
assert result["success"] is True
assert result["action"] == "deduct"
assert result["amount"] == 5.0
assert result["diff"] == 5.0
mock_deduct.assert_called_once()
dk = mock_deduct.call_args.kwargs
assert dk["amount"] == 5.0
assert dk["source"] == "viral_video"
assert dk["ref_id"] == "t4"
def test_deduct_insufficient_balance_returns_success_false_not_raise(self, service, db_session, user_id):
"""actual>estimated 补扣时余额不足(success=False)应记录 warning 但不抛异常。"""
import logging
deduct_res = {"success": False, "balance": 2.0, "transaction_id": None}
with (
patch.object(service, "deduct_points", return_value=deduct_res),
patch("packages.domain.points_service.logger") as mock_logger,
):
result = service.settle_viral_video(user_id, estimated=15.0, actual=20.0, txn_id="t5", db=db_session)
# 即使补扣失败,函数也返回 action=deduct 但 success=False(不阻塞任务完成)
assert result["success"] is False
assert result["action"] == "deduct"
assert result["amount"] == 5.0
# 应打印 warning
mock_logger.warning.assert_called_once()
def test_deduct_exception_returns_failure(self, service, db_session, user_id):
"""deduct_points 抛异常时应捕获并返回 success=False。"""
with patch.object(service, "deduct_points", side_effect=RuntimeError("db boom")):
result = service.settle_viral_video(user_id, estimated=10.0, actual=20.0, txn_id="t6", db=db_session)
assert result["success"] is False
assert result["action"] == "deduct"
assert result["diff"] == 10.0
class TestRefundViralVideo:
"""refund_viral_video(): 爆款视频失败全额退款。"""
def test_zero_amount_returns_none_action(self, service, db_session, user_id):
"""amount<=0 直接返回 none action,不调 refund_points。"""
with patch.object(service, "refund_points") as mock_refund:
r1 = service.refund_viral_video(user_id, 0, "t0", db_session)
r2 = service.refund_viral_video(user_id, None, "t0", db_session)
r3 = service.refund_viral_video(user_id, -1.5, "t0", db_session)
assert r1 == {"success": True, "action": "none", "amount": 0.0}
assert r2 == {"success": True, "action": "none", "amount": 0.0}
assert r3["action"] == "none"
mock_refund.assert_not_called()
def test_success_path_delegates_to_refund_points(self, service, db_session, user_id):
"""成功路径:透传 user_id/amount/ref_id=txn_id/source=viral_video。"""
expected = {"success": True, "balance": 80.0, "transaction_id": "rf-1"}
with patch.object(service, "refund_points", return_value=expected) as mock_refund:
result = service.refund_viral_video(user_id, 30.0, "txn-xyz", db_session)
assert result == expected
mock_refund.assert_called_once()
rk = mock_refund.call_args.kwargs
assert rk["user_id"] == user_id
assert rk["amount"] == 30.0
assert rk["source"] == "viral_video"
assert rk["ref_id"] == "txn-xyz"
assert rk["description"] == "爆款视频失败退款"
def test_exception_returns_failure(self, service, db_session, user_id):
"""refund_points 抛异常时返回 success=False/action=refund。"""
with patch.object(service, "refund_points", side_effect=RuntimeError("conn lost")):
result = service.refund_viral_video(user_id, 25.0, "txn-err", db_session)
assert result["success"] is False
assert result["action"] == "refund"
assert result["amount"] == 25.0
+2 -6
View File
@@ -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)
+11 -5
View File
@@ -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
+249 -1
View File
@@ -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