From 86248963799b969b37f9485dc90260952d000f79 Mon Sep 17 00:00:00 2001 From: Xiaoxia Agent Date: Mon, 5 Oct 2026 21:02:35 +0800 Subject: [PATCH 1/3] =?UTF-8?q?feat:=20=E5=8A=9F=E8=83=BD=E8=AE=A1?= =?UTF-8?q?=E8=B4=B9DB=E5=8C=96=EF=BC=88=E7=88=86=E6=AC=BE=E8=AF=BB?= =?UTF-8?q?=E9=85=8D=E7=BD=AE+=E5=AF=B9=E5=8F=A3=E5=9E=8B/=E6=99=BA?= =?UTF-8?q?=E8=83=BD=E5=89=AA=E8=BE=91=E8=AE=A1=E8=B4=B9=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../versions/096_feature_billing_fields.py | 61 +++ apps/api/app/api/routes/generation_tasks.py | 49 ++- apps/api/app/services/gpu_lipsync_service.py | 4 + apps/api/app/services/lipsync_service.py | 157 ++++++++ apps/api/app/tasks/lipsync_gpu.py | 29 ++ apps/worker/worker_app/tasks/generation.py | 48 +++ packages/adapters/sqlalchemy_impl/models.py | 14 + packages/domain/feature_pricing_service.py | 376 ++++++++++++++++++ packages/domain/points_rules.py | 91 ++++- .../unit/test_feature_billing_integration.py | 221 ++++++++++ tests/unit/test_feature_pricing_service.py | 235 +++++++++++ 11 files changed, 1270 insertions(+), 15 deletions(-) create mode 100755 alembic/versions/096_feature_billing_fields.py create mode 100755 packages/domain/feature_pricing_service.py create mode 100755 tests/unit/test_feature_billing_integration.py create mode 100755 tests/unit/test_feature_pricing_service.py diff --git a/alembic/versions/096_feature_billing_fields.py b/alembic/versions/096_feature_billing_fields.py new file mode 100755 index 000000000..434c8de28 --- /dev/null +++ b/alembic/versions/096_feature_billing_fields.py @@ -0,0 +1,61 @@ +"""功能计费积分字段(爆款/对口型/智能剪辑 DB 化计费)。 + +给 gpu_lipsync_tasks / generation_tasks / lipsync_jobs 三张表加积分字段: +- credits_prepaid: 提交任务时预扣积分 +- credits_cost: 最终结算积分 +- credits_transaction_id: 预扣流水 ID + +注意:feature_pricing_configs 配置表由 xiaoxia-admin 侧 migration 建立, +本仓库只读,不在此创建。 + +Revision ID: 096_feature_billing_fields +Revises: 095_viral_video_prompt_templates +Create Date: 2026-10-05 +""" + +import sqlalchemy as sa + +from alembic import op + +revision = "096_feature_billing_fields" +down_revision = "095_viral_video_prompt_templates" +branch_labels = None +depends_on = None + +_TABLES = ("gpu_lipsync_tasks", "generation_tasks", "lipsync_jobs") +_COLUMNS = ( + ("credits_prepaid", sa.Float(), "0"), + ("credits_cost", sa.Float(), "0"), + ("credits_transaction_id", sa.String(36), ""), +) + + +def _table_exists(conn, name: str) -> bool: + return name in sa.inspect(conn).get_table_names() + + +def upgrade() -> None: + conn = op.get_bind() + for table in _TABLES: + if not _table_exists(conn, table): + continue + existing = {c["name"] for c in sa.inspect(conn).get_columns(table)} + for col_name, col_type, default in _COLUMNS: + if col_name in existing: + continue + op.add_column( + table, + sa.Column(col_name, col_type, nullable=False, server_default=default), + ) + + +def downgrade() -> None: + conn = op.get_bind() + for table in _TABLES: + if not _table_exists(conn, table): + continue + existing = {c["name"] for c in sa.inspect(conn).get_columns(table)} + for col_name, _col_type, _default in _COLUMNS: + if col_name not in existing: + continue + op.drop_column(table, col_name) diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index 792ab424a..cad90e5ff 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -44,6 +44,7 @@ from packages.application import ( GetGenerationTaskUseCase, ListGeneratedVideosByTaskUseCase, ) +from packages.domain import feature_pricing_service from packages.domain.smart_match import smart_select_assets # #2035:文案关键词 → 素材分类 映射表(用于 smart_match category_match 维度) @@ -163,7 +164,6 @@ def _infer_expected_categories(script_tags: set[str] | None) -> set[str] | None: return matched or None - logger = logging.getLogger(__name__) router = APIRouter() @@ -700,6 +700,17 @@ def create_generation_task( logger.info("画中画已下线,strategy_id %s → one_take", effective_strategy_id) effective_strategy_id = "one_take" + # ── smart_edit 计费预扣(全局 points 开关 + 功能开关均开才扣) ── + # 首期固定价:dynamic_cost=0,price=(0+fixed_cost)×multiplier,price_cap 封顶。 + # 预览任务不扣费;按任务条数扣费,任一任务预扣失败(余额不足)整体拒绝。 + smart_edit_charge = 0.0 + charged_task_count = 0 + if not request.is_preview and feature_pricing_service.is_feature_enabled("smart_edit"): + unit_credits, _bd = feature_pricing_service.calculate_price("smart_edit", 0.0) + if unit_credits > 0: + smart_edit_charge = round(unit_credits * count, 2) + charged_task_count = count + # 批量生成(count>1):每个变体必须走与单视频完全相同的独立选片流程(#1743/#1749)。 # - 变体 0:clone 源 plan(不污染源 plan),变体 1..N-1 用 reselect_plan_for_variant # 完整重跑选片(素材级去重:fresh 优先 → 受控复用 overlap≤20% → 短素材禁复用); @@ -930,6 +941,42 @@ def create_generation_task( ) # 变体序号写入 extra_meta(响应/排查时可辨识) task.extra_meta["variant_index"] = task_index + + # smart_edit 逐条预扣(首期固定价,credits_cost=prepaid,不做结算) + task_txn_id = "" + if charged_task_count > 0: + from packages.domain.points_service import PointsService + + unit_credits = round(smart_edit_charge / count, 2) + res = PointsService().deduct_points( + user_id=user_id, + amount=unit_credits, + source="smart_edit", + db=db, + description="智能剪辑生成预扣", + ref_id=task.id, + ) + if not res.get("success"): + # 余额不足:退还本次请求已扣积分后整体拒绝 + already_charged = round(unit_credits * task_index, 2) + if already_charged > 0: + PointsService().refund_points( + user_id=user_id, + amount=already_charged, + source="smart_edit", + db=db, + ref_id=task.id, + description="智能剪辑批量提交失败退回", + ) + raise HTTPException( + status_code=402, + detail=(f"积分不足:智能剪辑每条需 {unit_credits:.2f} 积分,当前余额 {res.get('balance', 0)}"), + ) + task_txn_id = str(res.get("transaction_id") or "") + task.credits_prepaid = unit_credits + task.credits_cost = unit_credits + task.credits_transaction_id = task_txn_id + generation_task_repository.update(task) try: # 兜底关联编辑计划:前端未传 source_edit_plan_id 时, # 通过 template_id + user_id 在 DB 层直接查找最新的 plan。 diff --git a/apps/api/app/services/gpu_lipsync_service.py b/apps/api/app/services/gpu_lipsync_service.py index 7c8286231..0dc568a5a 100644 --- a/apps/api/app/services/gpu_lipsync_service.py +++ b/apps/api/app/services/gpu_lipsync_service.py @@ -228,6 +228,8 @@ class GpuLipsyncService: lipsync_job_id: str = "", user_id: str = "", project_id: str = "", + credits_prepaid: float = 0.0, + credits_transaction_id: str = "", ) -> GpuLipsyncTaskModel: task_id = str(uuid.uuid4()) now = datetime.now(UTC) @@ -240,6 +242,8 @@ class GpuLipsyncService: audio_url=audio_url, status="pending", attempt=0, + credits_prepaid=float(credits_prepaid or 0.0), + credits_transaction_id=str(credits_transaction_id or ""), created_at=now, updated_at=now, ) diff --git a/apps/api/app/services/lipsync_service.py b/apps/api/app/services/lipsync_service.py index f4a248514..13865bdc9 100644 --- a/apps/api/app/services/lipsync_service.py +++ b/apps/api/app/services/lipsync_service.py @@ -38,6 +38,7 @@ from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import LipsyncJobModel from packages.application.cosyvoice_service import CosyVoiceError from packages.config import get_api_settings +from packages.domain import feature_pricing_service from packages.domain.sentence_timings import ( compute_sentence_timings, probe_audio_duration, @@ -368,6 +369,8 @@ class LipsyncService: lipsync_job_id=job.id, user_id=job.user_id, project_id=job.project_id, + credits_prepaid=float(getattr(job, "credits_prepaid", 0) or 0), + credits_transaction_id=str(getattr(job, "credits_transaction_id", "") or ""), ) logger.info( "[lipsync] 已创建 GPU 任务(异步): job_id=%s gpu_task=%s", @@ -415,6 +418,121 @@ class LipsyncService: job.output_duration, ) + # ── lip_sync 计费辅助 ──────────────────────────────────────────────── + + @staticmethod + def _estimate_duration( + *, + audio_duration: Optional[float] = None, + sentence_timings: Optional[list] = None, + script_text: str = "", + ) -> float: + """预估音频/成片秒数。 + + 优先级:audio_duration(预合成前端已 ffprobe)> timings 末句 end_time > + 脚本字数 / 5 字每秒 > 默认 10 秒。 + """ + if audio_duration and float(audio_duration) > 0: + return float(audio_duration) + if sentence_timings: + max_end = 0.0 + for item in sentence_timings: + if isinstance(item, dict): + end = item.get("end_time") or item.get("end") or 0.0 + else: + end = 0.0 + try: + max_end = max(max_end, float(end)) + except (TypeError, ValueError): + continue + if max_end > 0: + return max_end + text = (script_text or "").strip() + if text: + return max(1.0, len(text) / 5.0) + return 10.0 + + def _settle_lip_sync(self, job: LipsyncJobModel, actual_duration: float) -> None: + """按实际时长结算(首期只退不补:final < prepaid 退差额,> 不补)。 + + 幂等:credits_cost 已 > 0 说明结算过,直接跳过。 + 结算失败不阻塞业务(结果已产出),仅记录日志。 + """ + try: + prepaid = float(getattr(job, "credits_prepaid", 0) or 0) + if prepaid <= 0: + return + if float(getattr(job, "credits_cost", 0) or 0) > 0: + return + feature_cfg = feature_pricing_service.get_feature_config("lip_sync") + unit_cost = float(feature_cfg.dynamic_unit_cost) if feature_cfg is not None else 0.0 + duration = float(actual_duration or 0.0) + if duration <= 0: + duration = self._estimate_duration( + sentence_timings=job.sentence_timings, + script_text=job.script_text, + ) + final_price, _bd = feature_pricing_service.calculate_price("lip_sync", duration * unit_cost) + final_price = round(float(final_price), 2) + job.credits_cost = final_price + if final_price < prepaid - 0.009: + refund = round(prepaid - final_price, 2) + from packages.domain.points_service import PointsService + + res = PointsService().refund_points( + user_id=job.user_id, + amount=refund, + source="lip_sync", + db=self.db, + ref_id=str(job.credits_transaction_id or job.id), + description="对口型结算退费", + ) + if not res.get("success"): + logger.warning( + "[lip_sync] 结算退费失败 job_id=%s refund=%.2f(不阻塞)", + job.id, + refund, + ) + # final > prepaid:首期只退不补,不补扣 + self.db.commit() + except Exception: # noqa: BLE001 + logger.exception("[lip_sync] 结算异常 job_id=%s(不阻塞结果)", job.id) + try: + self.db.rollback() + except Exception: # noqa: BLE001 + pass + + def _refund_lip_sync(self, job: LipsyncJobModel) -> None: + """任务失败/取消时全额退还预扣积分(credits_cost 已结算则退实际未消耗部分)。""" + try: + prepaid = float(getattr(job, "credits_prepaid", 0) or 0) + if prepaid <= 0: + return + txn_id = str(getattr(job, "credits_transaction_id", "") or "") + cost = float(getattr(job, "credits_cost", 0) or 0) + refund = round(prepaid - cost, 2) if cost > 0 else round(prepaid, 2) + if refund <= 0: + return + from packages.domain.points_service import PointsService + + res = PointsService().refund_points( + user_id=job.user_id, + amount=refund, + source="lip_sync", + db=self.db, + ref_id=txn_id or job.id, + description="对口型失败/取消退款", + ) + if res.get("success"): + job.credits_cost = prepaid # 标记已全额退回,防重复退 + self.db.commit() + except Exception: # noqa: BLE001 + logger.exception("[lip_sync] 退款异常 job_id=%s", job.id) + try: + self.db.rollback() + except Exception: # noqa: BLE001 + pass + # ── 创建任务 ────────────────────────────────────────────────────────── def create_job( @@ -466,6 +584,35 @@ class LipsyncService: if not isinstance(sentence_timings, list) or len(sentence_timings) == 0: raise MediaKitError("预合成模式 sentence_timings 不能为空", code="InvalidInput") + # 0.5 lip_sync 计费预扣(全局 points 开关 + 功能开关均开才扣) + prepaid_credits = 0.0 + prepaid_txn_id = "" + if feature_pricing_service.is_feature_enabled("lip_sync"): + est_duration = self._estimate_duration( + audio_duration=audio_duration, + sentence_timings=sentence_timings, + script_text=script_text, + ) + feature_cfg = feature_pricing_service.get_feature_config("lip_sync") + unit_cost = float(feature_cfg.dynamic_unit_cost) if feature_cfg is not None else 0.0 + dynamic_cost = est_duration * unit_cost + prepaid_credits, _bd = feature_pricing_service.calculate_price("lip_sync", dynamic_cost) + if prepaid_credits > 0: + from packages.domain.points_service import PointsService + + res = PointsService().deduct_points( + user_id=user_id, + amount=prepaid_credits, + source="lip_sync", + db=self.db, + description="对口型生成预扣", + ) + if not res.get("success"): + raise ValueError( + f"积分不足:本次对口型需 {prepaid_credits:.2f} 积分,当前余额 {res.get('balance', 0)}" + ) + prepaid_txn_id = str(res.get("transaction_id") or "") + # 1. 创建数据库记录 job_id = str(uuid.uuid4()) job = LipsyncJobModel( @@ -482,6 +629,8 @@ class LipsyncService: emotion=emotion or "", # 音频直传(含预合成)直接进入 pending(后续同步改为 submitted);TTS 模式进入 tts_processing status="tts_processing" if is_tts_mode else "pending", + credits_prepaid=prepaid_credits, + credits_transaction_id=prepaid_txn_id, ) self.db.add(job) self.db.flush() @@ -677,6 +826,8 @@ class LipsyncService: job.completed_at = _now job.updated_at = _now self.db.commit() + # lip_sync 超时全额退款 + self._refund_lip_sync(job) return job # 未提交的任务不轮询 @@ -702,6 +853,8 @@ class LipsyncService: job.completed_at = datetime.now(UTC) job.updated_at = datetime.now(UTC) self.db.commit() + # lip_sync 结算(只退不补) + self._settle_lip_sync(job, float(job.output_duration or 0.0)) # 异步转存自家 OSS try: from app.tasks.lipsync_tts import persist_output_video_task @@ -719,6 +872,8 @@ class LipsyncService: job.error_message = error.get("message", "任务执行失败") job.error_code = error.get("code", "TaskFailed") job.completed_at = datetime.now(UTC) + # lip_sync 失败全额退款(先退款再统一 commit) + self._refund_lip_sync(job) else: # 中间状态(running/processing/queued 等)同步到 DB,避免前端永远卡在 submitted if isinstance(mk_status, str) and mk_status: @@ -812,6 +967,8 @@ class LipsyncService: job.status = "cancelled" job.updated_at = datetime.now(UTC) self.db.commit() + # lip_sync 取消全额退款 + self._refund_lip_sync(job) self.db.refresh(job) return job diff --git a/apps/api/app/tasks/lipsync_gpu.py b/apps/api/app/tasks/lipsync_gpu.py index afbb89028..5480a49a2 100644 --- a/apps/api/app/tasks/lipsync_gpu.py +++ b/apps/api/app/tasks/lipsync_gpu.py @@ -104,6 +104,7 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str) job.updated_at = datetime.now(UTC) db.commit() logger.info("[lipsync_gpu_async] GPU 任务已被用户取消: job_id=%s", job_id) + _refund_lip_sync(db, job) return if final_task.status != "done": @@ -141,6 +142,7 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str) job_id, job.output_duration, ) + _settle_lip_sync(db, job, final_task) except Exception as exc: logger.exception("[lipsync_gpu_async] 异常: job_id=%s err=%s", job_id, exc) try: @@ -157,6 +159,33 @@ def lipsync_gpu_process_async(self, job_id: str, user_id: str, gpu_task_id: str) db.close() +def _settle_lip_sync(db: Session, job: LipsyncJobModel, gpu_task) -> None: + """GPU 成功后结算:同步 credits_cost 到 gpu 任务并按实际时长多退少不补。""" + try: + from app.services.lipsync_service import LipsyncService + + # GPU 任务表先同步结算结果(标记用) + LipsyncService._settle_lip_sync(job, float(getattr(gpu_task, "result_duration", 0) or 0.0)) + gpu_task.credits_cost = float(job.credits_cost or 0.0) + db.commit() + except Exception: # noqa: BLE001 + logger.exception("[lipsync_gpu_async] lip_sync 结算异常 job_id=%s(不阻塞)", job.id) + try: + db.rollback() + except Exception: # noqa: BLE001 + pass + + +def _refund_lip_sync(db: Session, job: LipsyncJobModel) -> None: + """GPU 取消/失败路径全额退款。""" + try: + from app.services.lipsync_service import LipsyncService + + LipsyncService(db)._refund_lip_sync(job) + except Exception: # noqa: BLE001 + logger.exception("[lipsync_gpu_async] lip_sync 退款异常 job_id=%s", job.id) + + def _fallback_to_mediakit(db: Session, job: LipsyncJobModel) -> None: """GPU 失败时回退到 MediaKit 云端渲染。""" try: diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 9a227c67f..00e2fb8d9 100644 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -387,6 +387,41 @@ BATCH_RENDER_SIMILARITY_LIMIT = 0.20 """批次内成片查重相似度阈值:超过则重选独立 plan 重渲一次(20%)。""" +def _refund_smart_edit_prepaid(task_id: str) -> None: + """智能剪辑任务最终失败时退还预扣积分(幂等)。""" + session = SessionLocal() + try: + from packages.adapters.sqlalchemy_impl.generation_task_repository import ( + SQLAlchemyGenerationTaskRepository, + ) + from packages.domain.points_service import PointsService + + repo = SQLAlchemyGenerationTaskRepository(session) + task = repo.get(task_id) + if not task: + return + prepaid = float(getattr(task, "credits_prepaid", 0) or 0) + if prepaid <= 0: + return + txn_id = getattr(task, "credits_transaction_id", "") or "" + res = PointsService().refund_points( + user_id=task.user_id, + amount=prepaid, + source="smart_edit", + db=session, + ref_id=task.id, + related_transaction_id=txn_id or None, + description="智能剪辑任务失败退回", + ) + task.credits_cost = 0.0 + task.credits_prepaid = 0.0 + repo.update(task) + if not res.get("success"): + logger.warning("[task_id=%s] 失败退积分未成功: %s", task_id, res) + finally: + session.close() + + def should_rerender_for_batch_dedup(*, batch_id: str, render_attempt: int, batch_similarity) -> bool: """批次内查重后判定是否需要重选 plan 重渲。 @@ -1167,6 +1202,10 @@ def generate_video(self, task_id: str) -> dict: "mark_failed", error_message="source_edit_plan_id is required. Please create a preview task first.", ) + try: + _refund_smart_edit_prepaid(task_id) + except Exception: + logger.warning("[task_id=%s] 失败退积分异常", task_id, exc_info=True) return { "status": "failed", "task_id": task_id, @@ -1205,6 +1244,7 @@ def generate_video(self, task_id: str) -> dict: ) # ── 自动重试逻辑 ────────────────────────────────────────────────── + will_retry = False try: from packages.adapters.sqlalchemy_impl.generation_task_repository import ( SQLAlchemyGenerationTaskRepository, @@ -1217,6 +1257,7 @@ def generate_video(self, task_id: str) -> dict: if _task and _task.auto_retry_enabled and _task.auto_retry_max > 0: current_retry = _task.retry_count or 0 if current_retry < _task.auto_retry_max: + will_retry = True logger.info( "[task_id=%s] 触发自动重试: 当前重试次数=%d, 最大重试次数=%d", task_id, @@ -1250,6 +1291,13 @@ def generate_video(self, task_id: str) -> dict: exc_info=True, ) + # 最终失败(不再重试):退还 smart_edit 预扣积分 + if not will_retry: + try: + _refund_smart_edit_prepaid(task_id) + except Exception: + logger.warning("[task_id=%s] 失败退积分异常", task_id, exc_info=True) + return { "status": "failed", "task_id": task_id, diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 7b3f164c3..457e911fb 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -335,6 +335,10 @@ class GenerationTaskModel(Base): bgm_config = Column(JSON, nullable=False, default=dict) extra_meta = Column("metadata", JSON, nullable=False, default=dict) logs = Column(Text, nullable=False, default="[]", server_default="[]") + # 功能计费(smart_edit):预扣积分 / 最终积分 / 预扣流水 ID + credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0") + credits_cost = Column(Float, nullable=False, default=0.0, server_default="0") + credits_transaction_id = Column(String(36), nullable=False, default="", server_default="") created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(UTC)) updated_at = Column( DateTime, @@ -727,6 +731,11 @@ class LipsyncJobModel(Base): # 精确句子时间戳(TTS 合成后由 silencedetect 计算,用于 B-roll 精确定位) sentence_timings = Column(JSON, nullable=True) # list[{index,text,start_time,end_time}] + # 功能计费(lip_sync):预扣积分 / 最终积分 / 预扣流水 ID + credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0") + credits_cost = Column(Float, nullable=False, default=0.0, server_default="0") + credits_transaction_id = Column(String(36), nullable=False, default="", server_default="") + # 时间戳 submitted_at = Column(DateTime, nullable=True) completed_at = Column(DateTime, nullable=True) @@ -905,6 +914,11 @@ class GpuLipsyncTaskModel(Base): # 心跳:worker 最近一次 poll/result 的时间,用于判定 worker 失联 last_heartbeat_at = Column(DateTime, nullable=True) + # 功能计费(lip_sync):预扣积分 / 最终积分 / 预扣流水 ID + credits_prepaid = Column(Float, nullable=False, default=0.0, server_default="0") + credits_cost = Column(Float, nullable=False, default=0.0, server_default="0") + credits_transaction_id = Column(String(36), nullable=False, default="", server_default="") + class GpuWorkerModel(Base): """GPU Worker 注册表 — 反向轮询模式下用于心跳与监控.""" diff --git a/packages/domain/feature_pricing_service.py b/packages/domain/feature_pricing_service.py new file mode 100755 index 000000000..f315187d5 --- /dev/null +++ b/packages/domain/feature_pricing_service.py @@ -0,0 +1,376 @@ +"""功能计费配置服务:从 feature_pricing_configs 读配置,300 秒 TTL 内存缓存。 + +配置表由 xiaoxia-admin 侧维护(同库 PostgreSQL),本服务只读。 +DB 不可用 / 表不存在 / 无数据时自动回落到内置兜底配置,保证业务不崩。 + +计费公式:最终积分 = (动态成本 + 固定成本) × 利润系数,price_cap 封顶。 +启用条件:全局 points_enabled 总开关 AND 功能 is_enabled 同时为 true。 +""" + +from __future__ import annotations + +import json +import logging +import threading +import time +from dataclasses import dataclass, field +from typing import Optional + +import sqlalchemy as sa + +from packages.adapters.sqlalchemy_impl import session as _session_mod + +logger = logging.getLogger(__name__) + +CACHE_TTL_SECONDS = 300.0 + +# ── 爆款视频兜底模型单价(与旧硬编码表/现状一致;DB 不可用时使用) ─────── +# 结构:models[model_key][resolution]["true"/"false"] = 单价 +# token 模式:元/百万输出 tokens;per_second 模式:元/秒 +# 注意:仅 seedance-2.5 配置 true(图生视频)单价;其余模型只有 false, +# 精确 key 缺失时由 points_rules 回落到 seedance-2.5/false(与旧现状一致)。 +_FALLBACK_VIRAL_MODEL_PRICING: dict = { + "seedance-2.5": { + "480p": {"false": 70.0, "true": 42.0}, + "720p": {"false": 70.0, "true": 42.0}, + "1080p": {"false": 77.0, "true": 46.0}, + }, + "seedance-2.0": { + "480p": {"false": 46.0}, + "720p": {"false": 46.0}, + "1080p": {"false": 51.0}, + "4k": {"false": 80.0}, + }, + "seedance-2.0-fast": { + "480p": {"false": 28.0}, + "720p": {"false": 28.0}, + }, + "seedance-2.0-mini": { + "480p": {"false": 9.2}, + "720p": {"false": 9.2}, + }, + "wan-3.0": { + "480p": {"false": 0.3}, + "720p": {"false": 0.6}, + "1080p": {"false": 1.2}, + }, +} + + +@dataclass +class FeatureConfig: + """功能计费配置快照。""" + + feature_key: str + name: str = "" + emoji: str = "" + is_enabled: bool = False + fixed_cost: float = 0.0 + profit_multiplier: float = 1.0 + dynamic_unit_cost: float = 0.0 + billing_mode: str = "model_based" + price_cap: float = 0.0 + model_pricing: dict = field(default_factory=dict) + description: str = "" + + +# ── 进程内缓存:(loaded_monotonic, {feature_key: FeatureConfig}) ────────── +_lock = threading.Lock() +_cache: Optional[tuple[float, dict[str, FeatureConfig]]] = None + + +def _fallback_configs() -> dict[str, FeatureConfig]: + """内置兜底配置:爆款启用(与现状一致),其余两个关闭。""" + return { + "viral_video": FeatureConfig( + feature_key="viral_video", + name="爆款视频", + emoji="🎬", + is_enabled=True, + fixed_cost=0.15, + profit_multiplier=1.3, + dynamic_unit_cost=0.0, + billing_mode="model_based", + price_cap=0.0, + model_pricing=json.loads(json.dumps(_FALLBACK_VIRAL_MODEL_PRICING)), + description="爆款视频动态定价(兜底配置)", + ), + "lip_sync": FeatureConfig( + feature_key="lip_sync", + name="对口型", + emoji="🎙️", + is_enabled=False, + fixed_cost=0.0, + profit_multiplier=1.0, + dynamic_unit_cost=0.0, + billing_mode="per_second", + price_cap=0.0, + description="对口型计费(兜底配置,默认关闭)", + ), + "smart_edit": FeatureConfig( + feature_key="smart_edit", + name="智能剪辑", + emoji="✂️", + is_enabled=False, + fixed_cost=0.0, + profit_multiplier=1.0, + dynamic_unit_cost=0.0, + billing_mode="model_based", + price_cap=0.0, + description="智能剪辑固定价计费(兜底配置,默认关闭)", + ), + } + + +_lazy_session = None + + +def _get_session(): + """优先用全局 SessionLocal(worker);否则按应用配置懒建同步引擎(api)。""" + global _lazy_session + if _session_mod.SessionLocal is not None: + return _session_mod.SessionLocal() + if _lazy_session is not None: + return _lazy_session() + try: + from packages.config import get_shared_settings + + url = str(get_shared_settings().database_url) + except Exception: # noqa: BLE001 + return None + if not url: + return None + url = url.replace("postgresql+asyncpg://", "postgresql+psycopg://") + if url.startswith("postgresql://"): + url = url.replace("postgresql://", "postgresql+psycopg://") + engine = sa.create_engine(url, pool_pre_ping=True, pool_size=2, max_overflow=2) + from sqlalchemy.orm import sessionmaker + + _lazy_session = sessionmaker(bind=engine) + return _lazy_session() + + +def _parse_model_pricing(raw) -> dict: + """解析 model_pricing_json(Text JSON),空/失败 → {}。""" + if raw is None: + return {} + if isinstance(raw, dict): + return raw + text = str(raw).strip() + if not text: + return {} + try: + data = json.loads(text) + except (ValueError, TypeError): + logger.warning("model_pricing_json 解析失败,按空配置处理: %r", text[:200]) + return {} + return data if isinstance(data, dict) else {} + + +def _to_float(value, default: float = 0.0) -> float: + try: + if value is None: + return default + return float(value) + except (TypeError, ValueError): + return default + + +def _load_all() -> dict[str, FeatureConfig]: + """SELECT * FROM feature_pricing_configs,返回 {feature_key: FeatureConfig}。 + + 表不存在 / DB 异常由调用方捕获并回落兜底配置。 + """ + session = None + try: + session = _get_session() + if session is None: + raise RuntimeError("no db session available") + sql = sa.text(""" + SELECT feature_key, name, emoji, is_enabled, fixed_cost, + profit_multiplier, dynamic_unit_cost, billing_mode, + price_cap, model_pricing_json, description + FROM feature_pricing_configs + """) + rows = session.execute(sql).mappings().all() + configs: dict[str, FeatureConfig] = {} + for row in rows: + key = str(row["feature_key"] or "").strip() + if not key: + continue + configs[key] = FeatureConfig( + feature_key=key, + name=str(row["name"] or key), + emoji=str(row["emoji"] or ""), + is_enabled=bool(row["is_enabled"]), + fixed_cost=_to_float(row["fixed_cost"]), + profit_multiplier=_to_float(row["profit_multiplier"], 1.0), + dynamic_unit_cost=_to_float(row["dynamic_unit_cost"]), + billing_mode=str(row["billing_mode"] or "model_based"), + price_cap=_to_float(row["price_cap"]), + model_pricing=_parse_model_pricing(row["model_pricing_json"]), + description=str(row["description"] or ""), + ) + return configs + finally: + if session is not None: + try: + session.close() + except Exception: # noqa: BLE001 + pass + + +def _get_cache() -> dict[str, FeatureConfig]: + """TTL 内返回缓存,否则重新 load;DB 异常/表不存在时返回内置兜底配置。""" + global _cache + now = time.monotonic() + with _lock: + if _cache is not None and now - _cache[0] < CACHE_TTL_SECONDS: + return _cache[1] + + try: + loaded = _load_all() + except Exception: # noqa: BLE001 - 表不存在/DB 不可用时静默回落 + logger.info("feature_pricing_configs 读取失败,使用内置兜底配置", exc_info=True) + return _fallback_configs() + + # DB 可用但表为空:同样回落兜底(保证爆款现状不被改变) + if not loaded: + fallback = _fallback_configs() + with _lock: + _cache = (now, fallback) + return fallback + + # 以兜底为底(DB 未配置的 feature_key 仍有兜底),DB 行覆盖 + merged = _fallback_configs() + merged.update(loaded) + with _lock: + _cache = (now, merged) + return merged + + +def get_feature_config(feature_key: str) -> Optional[FeatureConfig]: + """获取指定功能配置,未知 key 返回 None。""" + key = str(feature_key or "").strip() + if not key: + return None + return _get_cache().get(key) + + +def _global_points_enabled() -> bool: + """全局积分总开关(兼容 api / worker 运行时),取不到时默认关闭。""" + try: + from packages.shared import get_shared_settings + + return bool(get_shared_settings().points_enabled) + except Exception: # noqa: BLE001 + pass + try: + from app.config import settings + + return bool(getattr(settings, "points_enabled", False)) + except Exception: # noqa: BLE001 + return False + + +def is_feature_enabled(feature_key: str) -> bool: + """功能是否启用并扣费:全局 points_enabled AND 功能 is_enabled。""" + cfg = get_feature_config(feature_key) + if cfg is None: + return False + return bool(cfg.is_enabled) and _global_points_enabled() + + +def calculate_price(feature_key: str, dynamic_cost: float = 0.0) -> tuple[float, dict]: + """按公式计算最终积分并返回明细。 + + price = (dynamic_cost + fixed_cost) × profit_multiplier + price_cap > 0 时封顶(取 min)。 + 功能未启用 → (0.0, breakdown{is_enabled: False, charged: False})。 + """ + cfg = get_feature_config(feature_key) + dynamic = max(0.0, _to_float(dynamic_cost)) + if cfg is None or not cfg.is_enabled: + return 0.0, { + "feature_key": feature_key, + "is_enabled": False, + "charged": False, + "dynamic_cost": dynamic, + "fixed_cost": 0.0, + "profit_multiplier": 1.0, + "price_cap": 0.0, + "final_price": 0.0, + } + + fixed = max(0.0, cfg.fixed_cost) + multiplier = cfg.profit_multiplier if cfg.profit_multiplier > 0 else 1.0 + raw_price = (dynamic + fixed) * multiplier + cap = cfg.price_cap if cfg.price_cap and cfg.price_cap > 0 else 0.0 + final_price = min(raw_price, cap) if cap else raw_price + final_price = round(float(final_price), 2) + breakdown = { + "feature_key": cfg.feature_key, + "is_enabled": True, + "charged": True, + "dynamic_cost": round(dynamic, 4), + "fixed_cost": float(fixed), + "profit_multiplier": float(multiplier), + "price_cap": float(cap), + "raw_price": round(float(raw_price), 4), + "final_price": final_price, + } + return final_price, breakdown + + +def lookup_model_price( + model_pricing: dict, + model_key: str, + resolution: str, + has_video_input: bool, +) -> Optional[float]: + """从 model_pricing dict 取模型单价,兼容两种常见 JSON 结构。 + + 1. 嵌套:{model: {resolution: {"true"/"false": price}}} + (内层 bool key 也兼容直接 bool / 省略) + 2. 扁平:{"model|resolution|true_or_false": price} + (分隔符支持 | / : / , / 空格;bool 段可省略) + 取不到返回 None。 + """ + if not isinstance(model_pricing, dict): + return None + model = str(model_key or "").strip() + res = str(resolution or "").strip() + flag = "true" if has_video_input else "false" + + # 1. 嵌套 + model_node = model_pricing.get(model) + if isinstance(model_node, dict): + res_node = model_node.get(res) + if isinstance(res_node, dict): + # 精确 bool key 命中才返回;不做“只有一个值就取”的模糊匹配 + # (否则缺失 true 时会错误地取到 false 价,破坏旧版回落规则) + if flag in res_node: + return _to_float(res_node[flag]) if res_node[flag] is not None else None + if has_video_input in res_node: + val = res_node[has_video_input] + return _to_float(val) if val is not None else None + elif isinstance(res_node, (int, float)): + return float(res_node) + + # 2. 扁平 + for sep in ("|", ":", ",", " "): + for key in ( + f"{model}{sep}{res}{sep}{flag}", + f"{model}{sep}{res}", + ): + if key in model_pricing: + value = model_pricing[key] + return _to_float(value) if value is not None else None + return None + + +def refresh_feature_configs() -> None: + """清空缓存(下次读取重新 load DB;测试/admin 改配置后可手动调)。""" + global _cache + with _lock: + _cache = None diff --git a/packages/domain/points_rules.py b/packages/domain/points_rules.py index d226eb3ab..d03957bdf 100644 --- a/packages/domain/points_rules.py +++ b/packages/domain/points_rules.py @@ -2,17 +2,21 @@ v1.6.1: 按产品决策,智能混剪/AI数字人/AI配音/抖音解析/改写/标题/封面 全部免费, 仅保留声音克隆合成(voice_clone_synth)的扣点逻辑;声音克隆训练保持免费。 -爆款视频(viral_video)走动态定价,见本文件 VIRAL_VIDEO_MODEL_PRICES + calculate_viral_video_credits。 +爆款视频(viral_video)走动态定价,计费参数 DB 化(feature_pricing_configs, +见 feature_pricing_service),calculate_viral_video_credits 从配置读取单价/ +固定成本/利润系数/封顶,DB 不可用时回落兜底配置。 """ from __future__ import annotations import math -# ============ 爆款视频动态定价 (#2151) ============ -# key = (model_id, resolution, has_video_input),单位: -# - billing_mode=token: 元/百万tokens(输出) -# - billing_mode=per_second: 元/秒(视频时长) +from packages.domain import feature_pricing_service + +# ============ 爆款视频动态定价 ============ +# 单价/固定成本/利润系数已 DB 化(feature_pricing_configs,feature_key=viral_video), +# 由 feature_pricing_service 读取(300s 缓存),DB 不可用时回落内置兜底配置。 +# 以下三个常量仅为向后兼容保留(旧引用方/兜底场景),值取自兜底配置。 VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = { ("seedance-2.5", "480p", False): 70.0, ("seedance-2.5", "720p", False): 70.0, @@ -33,9 +37,9 @@ VIRAL_VIDEO_MODEL_PRICES: dict[tuple[str, str, bool], float] = { ("wan-3.0", "1080p", False): 1.2, } -# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器 +# 固定成本(元):VLM 分析 + LLM 文案 + TTS + OSS + 服务器(兜底默认值) VIRAL_VIDEO_FIXED_COST = 0.15 -# 利润系数 +# 利润系数(兜底默认值) VIRAL_VIDEO_PROFIT_MULTIPLIER = 1.3 # Seedance 输出帧率 VIRAL_VIDEO_FPS = 24 @@ -222,17 +226,22 @@ def calculate_viral_video_credits_with_breakdown( ) -> tuple[float, dict]: """计算爆款视频所需积分(1 积分 = 1 元),并返回计费公式明细。 + 单价/固定成本/利润系数/封顶从 feature_pricing_configs(viral_video)读取; + DB 不可用时回落与现状一致的内置兜底配置。 + 公式: tokens = duration * width * height * fps / 1024 video_cost = tokens / 1_000_000 * model_token_price total = round((video_cost + fixed_cost) * profit_multiplier, 2) + price_cap > 0 时封顶取 min 若传入 actual_tokens 则用它替代计算值。 Returns: (credits, breakdown) 二元组: - credits: 四舍五入保留两位小数的最终积分 - breakdown: dict,包含 tokens / video_cost / fixed_cost / profit_multiplier / - model_price / width / height / fps 字段,便于前端展示计费明细。 + model_price / width / height / fps / feature_enabled / charged / price_cap + 字段,便于前端展示计费明细。功能关闭时 credits=0、charged=False。 """ w = max(1, int(width or 1)) h = max(1, int(height or 1)) @@ -242,11 +251,36 @@ def calculate_viral_video_credits_with_breakdown( cfg = get_viral_video_model_config(prefix) res_key = _infer_resolution_key(w, h) billing = cfg.get("billing_mode", "token") - key = (prefix, res_key, bool(has_video_input)) - price = VIRAL_VIDEO_MODEL_PRICES.get(key) + dur = max(1, int(duration_seconds or 15)) + + # ── 从 DB 配置(兜底内置)取计费参数 ── + feature_cfg = feature_pricing_service.get_feature_config("viral_video") + # 注意:此处 feature_enabled 只表示“功能自身开关”,不并入全局 points_enabled + # 总开关(保持与旧版计费函数行为一致:价格照常计算)。全局总开关由业务层 + # (route/worker)通过 feature_pricing_service.is_feature_enabled 统一把关。 + feature_enabled = bool(feature_cfg.is_enabled) if feature_cfg is not None else True + model_pricing = feature_cfg.model_pricing if feature_cfg is not None else {} + fixed_cost = float(feature_cfg.fixed_cost) if feature_cfg is not None else float(VIRAL_VIDEO_FIXED_COST) + multiplier = ( + float(feature_cfg.profit_multiplier) + if feature_cfg is not None and feature_cfg.profit_multiplier > 0 + else float(VIRAL_VIDEO_PROFIT_MULTIPLIER) + ) + price_cap = float(feature_cfg.price_cap) if feature_cfg is not None else 0.0 + + # 单价:优先配置 dict;复刻旧版回落规则——精确 key 取不到时,回落 + # seedance-2.5 同分辨率 False 单价;最终兜底 70.0。 + price = feature_pricing_service.lookup_model_price(model_pricing, prefix, res_key, bool(has_video_input)) + if price is None: + # 配置表未命中:先尝试配置里的 seedance-2.5/False + if prefix != "seedance-2.5" or bool(has_video_input): + price = feature_pricing_service.lookup_model_price(model_pricing, "seedance-2.5", res_key, False) + if price is None: + key = (prefix, res_key, bool(has_video_input)) + price = VIRAL_VIDEO_MODEL_PRICES.get(key) if price is None: price = VIRAL_VIDEO_MODEL_PRICES.get(("seedance-2.5", res_key, False), 70.0) - dur = max(1, int(duration_seconds or 15)) + if billing == "per_second": tokens = 0.0 video_cost = dur * float(price) @@ -259,13 +293,40 @@ def calculate_viral_video_credits_with_breakdown( video_cost = tokens / 1_000_000.0 * float(price) billing_unit = "token" - total = (video_cost + VIRAL_VIDEO_FIXED_COST) * VIRAL_VIDEO_PROFIT_MULTIPLIER + if not feature_enabled: + # 功能关闭(is_enabled=false 或全局 points 关闭):不扣费,明细照旧返回 + credits = 0.0 + raw_total = (video_cost + fixed_cost) * multiplier + breakdown = { + "tokens": float(tokens), + "video_cost": float(video_cost), + "fixed_cost": float(fixed_cost), + "profit_multiplier": float(multiplier), + "price_cap": float(price_cap or 0.0), + "model_price": float(price), + "model_key": prefix, + "billing_mode": billing, + "billing_unit": billing_unit, + "width": int(w), + "height": int(h), + "fps": int(effective_fps), + "duration": dur, + "feature_enabled": False, + "charged": False, + "raw_price": round(float(raw_total), 4), + } + return credits, breakdown + + total = (video_cost + fixed_cost) * multiplier + if price_cap and price_cap > 0: + total = min(total, price_cap) credits = round(float(total), 2) breakdown = { "tokens": float(tokens), "video_cost": float(video_cost), - "fixed_cost": float(VIRAL_VIDEO_FIXED_COST), - "profit_multiplier": float(VIRAL_VIDEO_PROFIT_MULTIPLIER), + "fixed_cost": float(fixed_cost), + "profit_multiplier": float(multiplier), + "price_cap": float(price_cap or 0.0), "model_price": float(price), "model_key": prefix, "billing_mode": billing, @@ -274,6 +335,8 @@ def calculate_viral_video_credits_with_breakdown( "height": int(h), "fps": int(effective_fps), "duration": dur, + "feature_enabled": True, + "charged": True, } return credits, breakdown diff --git a/tests/unit/test_feature_billing_integration.py b/tests/unit/test_feature_billing_integration.py new file mode 100755 index 000000000..0d8690206 --- /dev/null +++ b/tests/unit/test_feature_billing_integration.py @@ -0,0 +1,221 @@ +"""功能计费改造测试:爆款读配置、对口型/智能剪辑预扣逻辑。 + +策略: +- 爆款:通过修改缓存中的 FeatureConfig(multiplier/model_pricing)验证价格随配置变化 +- lip_sync / smart_edit:直接测 LipsyncService 的预扣/结算/退款辅助方法, + PointsService 用 mock,避免依赖真实积分账户。 +""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +from packages.domain import feature_pricing_service as fps +from packages.domain.feature_pricing_service import FeatureConfig, refresh_feature_configs + + +@pytest.fixture(autouse=True) +def _reset_cache(): + refresh_feature_configs() + yield + refresh_feature_configs() + + +def _seed_cache(configs: dict) -> None: + import time + + fps._cache = (time.monotonic(), configs) + + +class TestViralVideoReadsConfig: + def test_multiplier_change_changes_price(self): + """配置里 multiplier 改大后,爆款价格随之变大(证明不再读死常量)。""" + from packages.domain.points_rules import calculate_viral_video_credits + + # 基线兜底 + base = calculate_viral_video_credits(15, 1280, 720) + assert base == 29.68 + + fallback = fps._fallback_configs() + vv = fallback["viral_video"] + vv.profit_multiplier = 2.0 + _seed_cache(fallback) + + changed = calculate_viral_video_credits(15, 1280, 720) + assert changed > base + # 精确校验:video_cost 相同,仅系数从 1.3 → 2.0 + _, bd = __import__( + "packages.domain.points_rules", fromlist=["calculate_viral_video_credits_with_breakdown"] + ).calculate_viral_video_credits_with_breakdown(15, 1280, 720) + assert bd["profit_multiplier"] == 2.0 + + def test_model_price_from_config(self): + """model_pricing 改单价后,token 成本按新单价计算。""" + from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown + + fallback = fps._fallback_configs() + vv = fallback["viral_video"] + # seedance-2.5/720p/false 从 70 改成 100 + vv.model_pricing["seedance-2.5"]["720p"]["false"] = 100.0 + _seed_cache(fallback) + + _, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720) + assert bd["model_price"] == 100.0 + + def test_price_cap_from_config(self): + from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown + + fallback = fps._fallback_configs() + vv = fallback["viral_video"] + vv.price_cap = 5.0 + _seed_cache(fallback) + + credits, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720) + assert credits == 5.0 + assert bd["price_cap"] == 5.0 + + def test_disabled_feature_returns_zero_credits(self): + """功能 is_enabled=false 时计费函数返回 0(纯计费层语义)。""" + from packages.domain.points_rules import calculate_viral_video_credits_with_breakdown + + fallback = fps._fallback_configs() + fallback["viral_video"].is_enabled = False + _seed_cache(fallback) + + credits, bd = calculate_viral_video_credits_with_breakdown(15, 1280, 720) + assert credits == 0.0 + assert bd["feature_enabled"] is False + assert bd["charged"] is False + + +class TestLipSyncPricing: + def _make_service(self): + from app.services.lipsync_service import LipsyncService + + svc = LipsyncService.__new__(LipsyncService) + svc.db = MagicMock() + return svc + + def _lip_cfg(self, **kw): + base = dict( + feature_key="lip_sync", + name="对口型", + is_enabled=True, + fixed_cost=0.1, + profit_multiplier=1.0, + dynamic_unit_cost=0.05, + billing_mode="per_second", + price_cap=0.0, + model_pricing={}, + description="", + ) + base.update(kw) + return FeatureConfig(**base) + + def test_estimate_duration_from_script(self): + svc = self._make_service() + # 10 个字 / 5 = 2 秒,下限 1 + assert svc._estimate_duration(script_text="一二三四五六七八九十") == 2.0 + # 无任何信息 → 默认 10 秒 + assert svc._estimate_duration() == 10.0 + + def test_calculate_lipsync_price_per_second(self): + _seed_cache({"lip_sync": self._lip_cfg()}) + price, bd = fps.calculate_price("lip_sync", dynamic_cost=20.0 * 0.05) + # dynamic 1.0 + fixed 0.1 = 1.1 + assert price == 1.1 + assert bd["charged"] is True + + def test_settle_refunds_overcharge(self): + """实际时长短 → 只退不补,退还差额。""" + svc = self._make_service() + _seed_cache({"lip_sync": self._lip_cfg()}) + + job = MagicMock() + job.credits_prepaid = 2.0 + job.credits_cost = 0.0 # 未结算 + job.user_id = "u1" + job.credits_transaction_id = "txn-old" + + with patch("packages.domain.points_service.PointsService") as MockPS: + inst = MockPS.return_value + inst.refund_points.return_value = {"success": True} + svc._settle_lip_sync(job, actual_duration=10.0) + + # final: (10*0.05 + 0.1)*1.0 = 0.6;退 2.0-0.6=1.4 + assert round(job.credits_cost, 2) == 0.6 + inst.refund_points.assert_called_once() + kwargs = inst.refund_points.call_args.kwargs + assert kwargs["amount"] == 1.4 + + def test_settle_no_refund_when_longer(self): + """首期只退不补:实际更贵不补扣。""" + svc = self._make_service() + _seed_cache({"lip_sync": self._lip_cfg()}) + + job = MagicMock() + job.credits_prepaid = 0.5 + job.credits_cost = 0.0 + + with patch("packages.domain.points_service.PointsService") as MockPS: + inst = MockPS.return_value + svc._settle_lip_sync(job, actual_duration=60.0) + + assert round(job.credits_cost, 2) > 0.5 + inst.refund_points.assert_not_called() + + def test_refund_on_failure_full(self): + svc = self._make_service() + job = MagicMock() + job.credits_prepaid = 3.0 + job.credits_cost = 0.0 + job.user_id = "u1" + job.credits_transaction_id = "t1" + + with patch("packages.domain.points_service.PointsService") as MockPS: + inst = MockPS.return_value + inst.refund_points.return_value = {"success": True} + svc._refund_lip_sync(job) + + kwargs = inst.refund_points.call_args.kwargs + assert kwargs["amount"] == 3.0 + + +class TestSmartEditFixedPrice: + def test_fixed_price_formula(self): + """首期固定价:dynamic=0,price=fixed*multiplier,cap 封顶。""" + cfg = FeatureConfig( + feature_key="smart_edit", + name="智能剪辑", + is_enabled=True, + fixed_cost=2.0, + profit_multiplier=1.5, + billing_mode="model_based", + price_cap=0.0, + ) + _seed_cache({"smart_edit": cfg}) + price, bd = fps.calculate_price("smart_edit", dynamic_cost=0.0) + # (0+2)*1.5 = 3.0 + assert price == 3.0 + assert bd["dynamic_cost"] == 0.0 + + def test_fixed_price_with_cap(self): + cfg = FeatureConfig( + feature_key="smart_edit", + is_enabled=True, + fixed_cost=10.0, + profit_multiplier=2.0, + price_cap=8.0, + ) + _seed_cache({"smart_edit": cfg}) + price, _ = fps.calculate_price("smart_edit", dynamic_cost=0.0) + assert price == 8.0 + + def test_disabled_smart_edit_free(self): + cfg = FeatureConfig(feature_key="smart_edit", is_enabled=False, fixed_cost=2.0) + _seed_cache({"smart_edit": cfg}) + price, bd = fps.calculate_price("smart_edit", dynamic_cost=0.0) + assert price == 0.0 + assert bd["charged"] is False diff --git a/tests/unit/test_feature_pricing_service.py b/tests/unit/test_feature_pricing_service.py new file mode 100755 index 000000000..033c4742a --- /dev/null +++ b/tests/unit/test_feature_pricing_service.py @@ -0,0 +1,235 @@ +"""feature_pricing_service 单元测试。 + +覆盖: +- 300s TTL 内存缓存(命中不重复 load / 过期重新 load / refresh 强制刷新) +- calculate_price 公式 (dynamic+fixed)*multiplier、price_cap 封顶、round +- disabled / 未知 key 返回 0 +- DB 异常 / 空表 → 内置兜底配置(爆款启用且价格与现状一致) +- lookup_model_price 嵌套/扁平结构与旧版回落语义 +""" + +from __future__ import annotations + +import time + +import pytest + +from packages.domain import feature_pricing_service as fps +from packages.domain.feature_pricing_service import ( + CACHE_TTL_SECONDS, + FeatureConfig, + calculate_price, + get_feature_config, + is_feature_enabled, + lookup_model_price, + refresh_feature_configs, +) + + +@pytest.fixture(autouse=True) +def _reset_cache(): + """每个用例前后清空模块缓存,避免相互污染。""" + refresh_feature_configs() + yield + refresh_feature_configs() + + +def _cfg(key="x", **kw) -> FeatureConfig: + base = dict( + feature_key=key, + name=key, + is_enabled=True, + fixed_cost=0.2, + profit_multiplier=2.0, + dynamic_unit_cost=0.0, + billing_mode="per_second", + price_cap=0.0, + model_pricing={}, + description="", + ) + base.update(kw) + return FeatureConfig(**base) + + +class TestCacheTTL: + def test_cache_hit_avoids_reload(self, monkeypatch): + """TTL 内第二次读取不再调 _load_all。""" + calls = {"n": 0} + + def fake_load(): + calls["n"] += 1 + return {"x": _cfg()} + + monkeypatch.setattr(fps, "_load_all", fake_load) + get_feature_config("x") + get_feature_config("x") + get_feature_config("x") + assert calls["n"] == 1 + + def test_expired_cache_reloads(self, monkeypatch): + """超过 TTL 后重新 load。""" + calls = {"n": 0} + + def fake_load(): + calls["n"] += 1 + return {"x": _cfg()} + + monkeypatch.setattr(fps, "_load_all", fake_load) + get_feature_config("x") + assert calls["n"] == 1 + + # 把缓存时间戳回拨到 TTL 之前 + ts, data = fps._cache + fps._cache = (ts - CACHE_TTL_SECONDS - 1, data) + get_feature_config("x") + assert calls["n"] == 2 + + def test_refresh_forces_reload(self, monkeypatch): + calls = {"n": 0} + + def fake_load(): + calls["n"] += 1 + return {"x": _cfg()} + + monkeypatch.setattr(fps, "_load_all", fake_load) + get_feature_config("x") + refresh_feature_configs() + get_feature_config("x") + assert calls["n"] == 2 + + def test_ttl_constant_is_300(self): + assert CACHE_TTL_SECONDS == 300.0 + + +class TestCalculatePrice: + def test_basic_formula(self, monkeypatch): + # (dynamic 1.0 + fixed 0.2) * 2.0 = 2.4 + monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(dynamic_unit_cost=1.0)}) + price, bd = calculate_price("x", dynamic_cost=1.0) + assert price == 2.4 + assert bd["dynamic_cost"] == 1.0 + assert bd["fixed_cost"] == 0.2 + assert bd["profit_multiplier"] == 2.0 + assert bd["final_price"] == 2.4 + assert bd["charged"] is True + + def test_price_cap_clamps(self, monkeypatch): + # raw = (1+0.2)*2 = 2.4,cap=1.0 → 1.0 + monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(price_cap=1.0)}) + price, bd = calculate_price("x", dynamic_cost=1.0) + assert price == 1.0 + assert bd["price_cap"] == 1.0 + + def test_no_cap_keeps_raw(self, monkeypatch): + # cap=0 视为不封顶 + monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(price_cap=0.0)}) + price, _ = calculate_price("x", dynamic_cost=1.0) + assert price == 2.4 + + def test_rounded_two_decimals(self, monkeypatch): + monkeypatch.setattr( + fps, + "_load_all", + lambda: {"x": _cfg(fixed_cost=0.1, profit_multiplier=1.0)}, + ) + price, _ = calculate_price("x", dynamic_cost=1.0 / 3.0) + # 0.3333... + 0.1 = 0.4333 → 0.43 + assert price == 0.43 + + def test_negative_dynamic_treated_as_zero(self, monkeypatch): + monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg()}) + price, _ = calculate_price("x", dynamic_cost=-5.0) + # (0 + 0.2) * 2 = 0.4 + assert price == 0.4 + + def test_disabled_returns_zero(self, monkeypatch): + monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=False)}) + price, bd = calculate_price("x", dynamic_cost=1.0) + assert price == 0.0 + assert bd["is_enabled"] is False + assert bd["charged"] is False + + def test_unknown_key_returns_zero(self, monkeypatch): + monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg()}) + price, bd = calculate_price("nope", dynamic_cost=1.0) + assert price == 0.0 + assert bd["charged"] is False + + +class TestDBFailureFallback: + def test_load_exception_uses_fallback(self, monkeypatch): + def boom(): + raise RuntimeError("table does not exist") + + monkeypatch.setattr(fps, "_load_all", boom) + cfg = get_feature_config("viral_video") + assert cfg is not None + assert cfg.is_enabled is True + assert cfg.fixed_cost == 0.15 + assert cfg.profit_multiplier == 1.3 + + def test_empty_table_uses_fallback(self, monkeypatch): + monkeypatch.setattr(fps, "_load_all", lambda: {}) + assert get_feature_config("viral_video").is_enabled is True + assert get_feature_config("lip_sync").is_enabled is False + assert get_feature_config("smart_edit").is_enabled is False + + def test_fallback_viral_price_matches_current(self, monkeypatch): + """兜底爆款价格与旧硬编码现状一致:seedance-2.5/720p/false=70。""" + monkeypatch.setattr(fps, "_load_all", lambda: {}) + from packages.domain.points_rules import calculate_viral_video_credits + + # 默认全局开关关闭,但纯计费函数价格照常算 + assert calculate_viral_video_credits(15, 1280, 720) == 29.68 + + def test_db_row_overrides_fallback(self, monkeypatch): + monkeypatch.setattr( + fps, + "_load_all", + lambda: {"viral_video": _cfg("viral_video", fixed_cost=0.5, profit_multiplier=2.0, price_cap=50.0)}, + ) + cfg = get_feature_config("viral_video") + assert cfg.fixed_cost == 0.5 + assert cfg.profit_multiplier == 2.0 + assert cfg.price_cap == 50.0 + + +class TestIsFeatureEnabled: + def test_disabled_feature(self, monkeypatch): + monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=False)}) + assert is_feature_enabled("x") is False + + def test_global_switch_off_blocks_enabled_feature(self, monkeypatch): + monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=True)}) + monkeypatch.setattr(fps, "_global_points_enabled", lambda: False) + assert is_feature_enabled("x") is False + + def test_both_switches_on(self, monkeypatch): + monkeypatch.setattr(fps, "_load_all", lambda: {"x": _cfg(is_enabled=True)}) + monkeypatch.setattr(fps, "_global_points_enabled", lambda: True) + assert is_feature_enabled("x") is True + + +class TestLookupModelPrice: + NESTED = { + "seedance-2.5": { + "720p": {"false": 70.0, "true": 42.0}, + }, + "wan-3.0": {"480p": {"false": 0.3}}, + } + + def test_nested_exact_hit(self): + assert lookup_model_price(self.NESTED, "seedance-2.5", "720p", False) == 70.0 + assert lookup_model_price(self.NESTED, "seedance-2.5", "720p", True) == 42.0 + + def test_missing_bool_key_returns_none(self): + # wan-3.0/480p 只有 false,请求 true → None(由调用方回落) + assert lookup_model_price(self.NESTED, "wan-3.0", "480p", True) is None + + def test_unknown_model_returns_none(self): + assert lookup_model_price(self.NESTED, "nope", "720p", False) is None + + def test_flat_structure(self): + flat = {"m|720p|false": 12.5} + assert lookup_model_price(flat, "m", "720p", False) == 12.5 + assert lookup_model_price(flat, "m", "720p", True) is None -- 2.54.0 From d949e900513eb8b89fa67350c685cf290a1c0345 Mon Sep 17 00:00:00 2001 From: Xiaoxia Agent Date: Mon, 5 Oct 2026 21:29:04 +0800 Subject: [PATCH 2/3] =?UTF-8?q?fix(style):=20E741=20=E6=A8=A1=E7=B3=8A?= =?UTF-8?q?=E5=8F=98=E9=87=8F=E5=90=8D=20l=20->=20lpos=EF=BC=88vision=20JS?= =?UTF-8?q?ON=E5=AE=9A=E4=BD=8D=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- apps/worker/worker_app/tasks/vision/vlm_fallback.py | 6 +++--- apps/worker/worker_app/tasks/vision/vlm_fast_json.py | 6 +++--- 2 files changed, 6 insertions(+), 6 deletions(-) diff --git a/apps/worker/worker_app/tasks/vision/vlm_fallback.py b/apps/worker/worker_app/tasks/vision/vlm_fallback.py index c9b9e1d61..eef12b395 100644 --- a/apps/worker/worker_app/tasks/vision/vlm_fallback.py +++ b/apps/worker/worker_app/tasks/vision/vlm_fallback.py @@ -140,11 +140,11 @@ def call_pro_vlm( usage.get("completion_tokens", 0), ) text = _strip_code_fence(raw) - l, r_pos = text.find("{"), text.rfind("}") - if l < 0 or r_pos <= l: + lpos, r_pos = text.find("{"), text.rfind("}") + if lpos < 0 or r_pos <= lpos: logger.warning("[vision.v2] pro 无JSON elapsed=%.1fs head=%s", elapsed, raw[:200]) return None - obj = json.loads(text[l : r_pos + 1]) + obj = json.loads(text[lpos : r_pos + 1]) if not isinstance(obj, dict): return None scene = obj.get("scene") or "通用" diff --git a/apps/worker/worker_app/tasks/vision/vlm_fast_json.py b/apps/worker/worker_app/tasks/vision/vlm_fast_json.py index 3de9bba1b..48af71262 100644 --- a/apps/worker/worker_app/tasks/vision/vlm_fast_json.py +++ b/apps/worker/worker_app/tasks/vision/vlm_fast_json.py @@ -152,9 +152,9 @@ def call_fast_json( reasoning_tokens, ) text = _strip_code_fence(raw) - l, r = text.find("{"), text.rfind("}") - if l >= 0 and r > l: - text = text[l : r + 1] + lpos, r = text.find("{"), text.rfind("}") + if lpos >= 0 and r > lpos: + text = text[lpos : r + 1] try: obj = json.loads(text) except json.JSONDecodeError: -- 2.54.0 From e2de75c9f6bef6d8f2a6836702a5318ee2314688 Mon Sep 17 00:00:00 2001 From: Xiaoxia Agent Date: Mon, 5 Oct 2026 21:55:57 +0800 Subject: [PATCH 3/3] =?UTF-8?q?fix(test):=20=E9=80=82=E9=85=8D#2200/#2207?= =?UTF-8?q?=20V2=E5=9B=BE=E7=89=87=E5=88=86=E6=9E=90=E6=89=B9=E5=A4=84?= =?UTF-8?q?=E7=90=86=E6=9E=B6=E6=9E=84=EF=BC=88=E7=A7=BB=E9=99=A4=5Fanalyz?= =?UTF-8?q?e=5Fsingle=5Fimage=E6=96=AD=E8=A8=80=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/unit/test_viral_video_wiring.py | 69 ++++++++++++++++++++------- 1 file changed, 51 insertions(+), 18 deletions(-) diff --git a/tests/unit/test_viral_video_wiring.py b/tests/unit/test_viral_video_wiring.py index fe5ada00e..81d2bfb12 100644 --- a/tests/unit/test_viral_video_wiring.py +++ b/tests/unit/test_viral_video_wiring.py @@ -97,21 +97,43 @@ def invalidate_loader_cache(): class TestImageAnalysisWiring: - def test_uses_loader_template_and_xml_parse(self, job): + def test_step_image_analysis_uses_v2_batch_path(self, job): + """#2200/#2207 后图片分析走 V2 批处理(OCR+lite JSON 并行), + _step_image_analysis 归一化 URL 后调用 analyze_images_v2。""" from apps.worker.worker_app.tasks import viral_video as vv - with patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML) as mock_v: - result = vv._analyze_single_image(0, "https://img/1.jpg", "vlm-lite", 15) + fake_product = { + "name": "lipstick", + "brand": "品牌X", + "category": "唇部彩妆", + "key_features": ["显白", "持久"], + "text_on_package": ["品牌X", "211"], + "_source": "v2", + } + with patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw): + with patch( + "worker_app.tasks.vision.analyze_images_v2", + return_value=[fake_product, fake_product], + create=True, + ) as mock_v2: + result = vv._step_image_analysis(job) - mock_v.assert_called_once() - # 验证调用时传入了 system_prompt(说明走了 loader 渲染的模板) - call_kwargs = mock_v.call_args.kwargs - assert "system_prompt" in call_kwargs and call_kwargs["system_prompt"] - # 结果包含从 XML 解析出的产品信息 - assert result["name"] == "lipstick" - assert result["brand"] == "品牌X" - assert "显白" in result["key_features"] - assert result["text_on_package"] == ["品牌X", "211"] + mock_v2.assert_called_once() + # 传入的是归一化后的图片 URL 列表 + assert mock_v2.call_args.args[0] == job.images + products = result["products"] + assert len(products) == 2 + assert products[0]["name"] == "lipstick" + assert products[0]["brand"] == "品牌X" + assert "显白" in products[0]["key_features"] + assert products[0]["text_on_package"] == ["品牌X", "211"] + + def test_step_image_analysis_empty_images(self, job): + from apps.worker.worker_app.tasks import viral_video as vv + + job.images = [] + result = vv._step_image_analysis(job) + assert result == {"products": []} # ── 2) 意图解析走模板 ─────────────────────────────────────────────── @@ -252,18 +274,29 @@ class TestEndToEndLoaderUsed: called_types.append(prompt_type) return real_get(prompt_type, **kwargs) + v2_product = { + "name": "lipstick", + "brand": "品牌X", + "key_features": ["显白", "持久"], + } with ( patch.object(pl, "get_template", side_effect=spy_get), - patch("packages.shared.ai_service.call_vision", return_value=IMAGE_XML), + patch.object(vv, "_normalize_image_url", side_effect=lambda raw, idx: raw), + patch( + "worker_app.tasks.vision.analyze_images_v2", + return_value=[v2_product], + create=True, + ), patch("packages.shared.ai_service.call_llm", return_value=INTENT_XML), ): - # 1) image - img_res = vv._analyze_single_image(0, "https://img/1.jpg", "vlm", 15) - # 2) intent + # 1) image(V2 路径,不再经过 prompt_loader) + img_step = vv._step_image_analysis(job) + img_res = img_step["products"][0] + # 2) intent(走 loader image_analysis? 否——intent_parsing 模板) intent_res = vv._step_intent_parsing(job, {"products": [img_res]}) - # 前两步分别调用了 image_analysis 和 intent_parsing - assert "image_analysis" in called_types + # V2 图片分析不再调用 loader;意图解析调用 intent_parsing 模板 + assert "image_analysis" not in called_types assert "intent_parsing" in called_types # script 和 review 单独验证(需要不同的 LLM 返回) -- 2.54.0