diff --git a/apps/api/app/api/routes/ai_avatar_render.py b/apps/api/app/api/routes/ai_avatar_render.py index 63f12ef6d..52fc98abb 100644 --- a/apps/api/app/api/routes/ai_avatar_render.py +++ b/apps/api/app/api/routes/ai_avatar_render.py @@ -29,6 +29,8 @@ from app.services.ai_avatar_render_service import ( from fastapi import APIRouter, Depends, HTTPException, Query from sqlalchemy.orm import Session +from packages.middleware.points_gate import points_gate + logger = logging.getLogger(__name__) router = APIRouter() @@ -42,10 +44,12 @@ def _get_service(db: Session = Depends(get_db_session)) -> AiAvatarRenderService @router.post("", response_model=AiAvatarRenderJobResponse, status_code=201) +@points_gate("ai_digital_human", per_unit=15) def create_render_job( body: CreateAiAvatarRenderRequest, current_user: AuthenticatedUser = Depends(get_current_user), svc: AiAvatarRenderService = Depends(_get_service), + db: Session = Depends(get_db_session), ): """提交 AI 数字人渲染任务. diff --git a/apps/api/app/api/routes/generation_preview.py b/apps/api/app/api/routes/generation_preview.py index eb2302ce7..b518b16d4 100755 --- a/apps/api/app/api/routes/generation_preview.py +++ b/apps/api/app/api/routes/generation_preview.py @@ -43,6 +43,7 @@ from packages.application import ( GetGenerationTaskUseCase, ListGeneratedVideosByTaskUseCase, ) +from packages.middleware.points_gate import points_gate logger = logging.getLogger(__name__) @@ -271,6 +272,7 @@ def _variant_value(values: list[str], index: int, fallback: str = "") -> str: @router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, status_code=201) +@points_gate("ai_video", quantity_field="preview_count") def create_preview_generation_task( request: CreatePreviewGenerationTaskRequest, authenticated_user: AuthenticatedUser = Depends(get_current_user), diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index d1e301ae1..ac4ae0d6b 100755 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -42,6 +42,7 @@ from packages.application import ( ListGeneratedVideosByTaskUseCase, ) from packages.domain.smart_match import smart_select_assets +from packages.middleware.points_gate import points_gate logger = logging.getLogger(__name__) @@ -211,6 +212,7 @@ def _resolve_project_and_library( @router.post("/tasks", response_model=BatchGenerationTaskResponse) +@points_gate("ai_video", quantity_field="count") def create_generation_task( request: CreateGenerationTaskRequest, authenticated_user: AuthenticatedUser = Depends(get_current_user), diff --git a/apps/api/app/api/routes/lipsync.py b/apps/api/app/api/routes/lipsync.py index ff3d84b6b..d13ce50dd 100644 --- a/apps/api/app/api/routes/lipsync.py +++ b/apps/api/app/api/routes/lipsync.py @@ -12,9 +12,11 @@ from __future__ import annotations import logging +import math from datetime import UTC from app.auth import AuthenticatedUser, get_current_user +from app.config import settings from app.dependencies import ( get_db_session, get_voice_clone_profile_repository, @@ -30,6 +32,9 @@ from app.services.mediakit_client import MediaKitError from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query from sqlalchemy.orm import Session +from packages.domain.points_rules import calculate_points_cost +from packages.domain.points_service import PointsService + logger = logging.getLogger(__name__) router = APIRouter() @@ -53,8 +58,40 @@ def _get_service( def create_lipsync_job( body: CreateLipsyncJobRequest, current_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), svc: LipsyncService = Depends(_get_service), ): + user_id = current_user.user.id + + # ── 积分扣点(#1895 P2) ── + _points_deducted = 0 + _points_scene = "ai_digital_human" + _points_svc = PointsService() if settings.points_enabled else None + if _points_svc is not None: + # 口型同步:TTS 模式按 script_text 估时长(240字/分钟);音频直传按 audio_duration(秒→分钟) + if body.audio_url and body.audio_duration and body.audio_duration > 0: + est_minutes = max(1.0, math.ceil(body.audio_duration / 60.0)) + elif body.script_text: + est_minutes = max(1.0, math.ceil(len(body.script_text) / 240)) + else: + est_minutes = 1.0 + _points_deducted = calculate_points_cost( + _points_scene, + is_member=getattr(current_user.user, "is_member", False), + duration_minutes=est_minutes, + member_type=getattr(current_user.user, "member_type", None), + ) + _deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db) + if not _deduct_res["success"]: + raise HTTPException( + status_code=402, + detail={ + "code": "INSUFFICIENT_POINTS", + "message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}", + "required": _points_deducted, + "balance": _deduct_res["balance"], + }, + ) """提交对口型任务. 三种模式: @@ -66,7 +103,7 @@ def create_lipsync_job( """ try: job = svc.create_job( - user_id=current_user.user.id, + user_id=user_id, video_url=body.video_url, audio_url=body.audio_url, audio_duration=body.audio_duration, @@ -79,8 +116,18 @@ def create_lipsync_job( project_id=body.project_id, ) except ValueError as exc: + if _points_deducted > 0 and _points_svc is not None: + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db) + except Exception as refund_err: + logger.warning(f"对口型 ValueError 退积分异常: err={refund_err}") raise HTTPException(status_code=400, detail=str(exc)) from exc except MediaKitError as exc: + if _points_deducted > 0 and _points_svc is not None: + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db) + except Exception as refund_err: + logger.warning(f"对口型 MediaKitError 退积分异常: err={refund_err}") status_code = 502 if exc.code in ("VoiceForbidden",): status_code = 403 @@ -96,11 +143,24 @@ def create_lipsync_job( ) from exc except Exception as exc: logger.error("创建对口型任务异常: %s", exc, exc_info=True) + if _points_deducted > 0 and _points_svc is not None: + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db) + except Exception as refund_err: + logger.warning(f"对口型异常退积分异常: err={refund_err}") raise HTTPException( status_code=400, detail=f"创建对口型任务失败: {exc}", ) from exc + # 创建成功但状态为 failed(同步路径失败已抛异常到上面 except;此处处理 Celery 调度失败等) + # 若任务已创建且状态为 failed,退费 + if _points_deducted > 0 and _points_svc is not None and getattr(job, "status", None) == "failed": + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id) + except Exception as refund_err: + logger.warning(f"对口型任务失败退积分异常: job_id={job.id}, err={refund_err}") + return job @@ -111,8 +171,34 @@ def create_lipsync_job( def preview_tts( body: AiAvatarTtsPreviewRequest, current_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), svc: LipsyncService = Depends(_get_service), ): + user_id = current_user.user.id + + # ── 积分扣点(#1895 P2) ── + _points_deducted = 0 + _points_scene = "ai_digital_human" + _points_svc = PointsService() if settings.points_enabled else None + if _points_svc is not None: + est_minutes = max(1.0, math.ceil(len(body.script_text or "") / 240)) if body.script_text else 1.0 + _points_deducted = calculate_points_cost( + _points_scene, + is_member=getattr(current_user.user, "is_member", False), + duration_minutes=est_minutes, + member_type=getattr(current_user.user, "member_type", None), + ) + _deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db) + if not _deduct_res["success"]: + raise HTTPException( + status_code=402, + detail={ + "code": "INSUFFICIENT_POINTS", + "message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}", + "required": _points_deducted, + "balance": _deduct_res["balance"], + }, + ) """步骤1「生成配音」同步 TTS 预合成. 同步执行 TTS 合成 → 下载音频 → ffprobe 时长 → 句子时间戳计算, @@ -121,13 +207,18 @@ def preview_tts( """ try: result = svc.preview_tts( - user_id=current_user.user.id, + user_id=user_id, voice_id=body.voice_id, script_text=body.script_text, speed=body.speed, emotion=body.emotion, ) except MediaKitError as exc: + if _points_deducted > 0 and _points_svc is not None: + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db) + except Exception as refund_err: + logger.warning(f"TTS 预合成 MediaKitError 退积分异常: err={refund_err}") status_code = 400 if exc.code in ("VoiceForbidden",): status_code = 403 @@ -142,6 +233,11 @@ def preview_tts( ) from exc except Exception as exc: logger.error("TTS 预合成异常: %s", exc, exc_info=True) + if _points_deducted > 0 and _points_svc is not None: + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db) + except Exception as refund_err: + logger.warning(f"TTS 预合成异常退积分异常: err={refund_err}") raise HTTPException( status_code=400, detail=f"TTS 合成失败: {exc}", diff --git a/apps/api/app/api/routes/scripts_ai.py b/apps/api/app/api/routes/scripts_ai.py index 495551e7a..78a999942 100644 --- a/apps/api/app/api/routes/scripts_ai.py +++ b/apps/api/app/api/routes/scripts_ai.py @@ -13,6 +13,7 @@ import re import tempfile from app.auth import AuthenticatedUser, get_current_user +from app.dependencies import get_db_session from app.schemas.scripts_ai import ( AiGenerateTitlesRequest, AiGenerateTitlesResponse, @@ -27,7 +28,9 @@ from app.services.script_asr_service import ( transcribe_to_text, ) from fastapi import APIRouter, Depends, HTTPException, status +from sqlalchemy.orm import Session +from packages.middleware.points_gate import points_gate from packages.shared.ai_client import get_doubao_client logger = logging.getLogger(__name__) @@ -62,9 +65,11 @@ def _validate_douyin_url(url: str) -> None: "/extract-from-douyin", response_model=ExtractFromDouyinResponse, ) +@points_gate("douyin_extract") def extract_from_douyin( request: ExtractFromDouyinRequest, - authenticated_user: AuthenticatedUser = Depends(get_current_user), + current_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), ) -> ExtractFromDouyinResponse: """从抖音视频下载无水印视频并通过 ASR 提取文案.""" source_url = request.url.strip() @@ -138,9 +143,11 @@ def extract_from_douyin( "/ai-rewrite", response_model=AiRewriteResponse, ) +@points_gate("ai_rewrite") def ai_rewrite( request: AiRewriteRequest, - authenticated_user: AuthenticatedUser = Depends(get_current_user), + current_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), ) -> AiRewriteResponse: """使用豆包大模型改写文案.""" content = (request.content or "").strip() @@ -206,9 +213,11 @@ def ai_rewrite( "/ai-generate-titles", response_model=AiGenerateTitlesResponse, ) +@points_gate("ai_title") def ai_generate_titles( request: AiGenerateTitlesRequest, - authenticated_user: AuthenticatedUser = Depends(get_current_user), + current_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), ) -> AiGenerateTitlesResponse: """使用现有 generate_smart_titles 生成标题.""" content = (request.content or "").strip() diff --git a/apps/api/app/api/routes/tts.py b/apps/api/app/api/routes/tts.py index 8a2e83b3a..a01c33d12 100644 --- a/apps/api/app/api/routes/tts.py +++ b/apps/api/app/api/routes/tts.py @@ -4,12 +4,14 @@ from __future__ import annotations import json import logging +import math import subprocess import tempfile from pathlib import Path from typing import Any, Optional from app.auth import AuthenticatedUser, get_current_user +from app.config import settings from app.core.celery_app import celery_app from app.core.storage import get_storage_service from app.dependencies import ( @@ -51,6 +53,8 @@ from packages.application.tts_job.use_cases import ( ) from packages.application.tts_job.workflow import TTSWorkflowService from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus +from packages.domain.points_rules import calculate_points_cost +from packages.domain.points_service import PointsService from packages.domain.voice_presets import list_voices from packages.ports.asset_library_repository import AssetLibraryRepository from packages.ports.asset_repository import AssetRepository @@ -128,6 +132,7 @@ def _to_response(job, sign_url=None) -> TTSJobResponse: def synthesize( request: TTSSynthesizeRequest, authenticated_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), repository: SQLAlchemyTTSJobRepository = Depends(_get_repository), cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service), voice_clone_repo=Depends(get_voice_clone_profile_repository), @@ -139,6 +144,31 @@ def synthesize( """ user_id = authenticated_user.user.id + # ── 积分扣点(#1895 P2) ── + _points_deducted = 0 + _points_scene = "ai_voice" + _points_svc = PointsService() if settings.points_enabled else None + if _points_svc is not None: + # 中文按 ~240 字/分钟粗估时长,至少按 1 分钟扣 1 分 + est_minutes = max(1.0, math.ceil(len(request.text) / 240)) + _points_deducted = calculate_points_cost( + _points_scene, + is_member=getattr(authenticated_user.user, "is_member", False), + duration_minutes=est_minutes, + member_type=getattr(authenticated_user.user, "member_type", None), + ) + _deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db) + if not _deduct_res["success"]: + raise HTTPException( + status_code=402, + detail={ + "code": "INSUFFICIENT_POINTS", + "message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}", + "required": _points_deducted, + "balance": _deduct_res["balance"], + }, + ) + # 解析 voice_id:前端可能传克隆音色 profile UUID(而非 CosyVoice voice_id), # 与 /tts/preview 保持一致:命中 profile → 校验归属 → 取 CosyVoice voice_id actual_voice_id = request.voice_id @@ -198,6 +228,7 @@ def synthesize( cosyvoice_service=cosyvoice_service, ) + synthesis_error: Exception | None = None try: job = workflow.start_synthesis(job.id) except Exception as e: @@ -205,10 +236,17 @@ def synthesize( # 但 DB 异常、网络异常等意外错误可能逃逸。 # 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。 logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True) + synthesis_error = e try: job = workflow.process_synthesis_failure(job.id, str(e)) except Exception as inner_e: logger.error(f"标记 TTS job 失败时出错: job_id={job.id}, error={inner_e}") + # 合成失败且已扣积分 → 退费 + if synthesis_error is not None and _points_deducted > 0 and _points_svc is not None: + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id) + except Exception as refund_err: + logger.warning(f"TTS 合失败退积分异常: job_id={job.id}, err={refund_err}") # 若任务处于 processing 状态(异步模式),触发 Celery 后台轮询 if job.status.value == "processing": @@ -223,10 +261,17 @@ def synthesize( celery_app.send_task("worker.process_tts_synthesis", args=[job.id]) except Exception as e: # Celery 调度失败,标记 job 为 failed + # e used below for refund context try: workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}") except Exception as inner_e: logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={inner_e}") + # 调度失败退费 + if _points_deducted > 0 and _points_svc is not None: + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db, ref_id=job.id) + except Exception as refund_err: + logger.warning(f"Celery 调度失败退积分异常: job_id={job.id}, err={refund_err}") return TTSSynthesizeResponse( job_id=job.id, @@ -553,6 +598,7 @@ def save_tts_job_to_library( def preview_tts( request: TTSPreviewRequest, authenticated_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service), voice_clone_repo=Depends(get_voice_clone_profile_repository), ) -> TTSPreviewResponse: @@ -561,6 +607,31 @@ def preview_tts( 用于前端预览配音效果,限制文本长度 200 字以内。 支持预设音色和克隆音色:克隆音色传的是 profile UUID,需解析为 CosyVoice voice_id。 """ + user_id = authenticated_user.user.id + # ── 积分扣点(#1895 P2) ── + _points_deducted = 0 + _points_scene = "ai_voice" + _points_svc = PointsService() if settings.points_enabled else None + if _points_svc is not None: + est_minutes = max(1.0, math.ceil(len(request.text) / 240)) + _points_deducted = calculate_points_cost( + _points_scene, + is_member=getattr(authenticated_user.user, "is_member", False), + duration_minutes=est_minutes, + member_type=getattr(authenticated_user.user, "member_type", None), + ) + _deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db) + if not _deduct_res["success"]: + raise HTTPException( + status_code=402, + detail={ + "code": "INSUFFICIENT_POINTS", + "message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}", + "required": _points_deducted, + "balance": _deduct_res["balance"], + }, + ) + # 解析 voice_id:前端可能传 VoiceCloneProfile UUID 或预设音色 ID actual_voice_id = request.voice_id profile = voice_clone_repo.get(request.voice_id) @@ -586,16 +657,16 @@ def preview_tts( emotion=request.emotion, language=getattr(request, "language", "zh-CN"), ) - except CosyVoiceError as e: - raise HTTPException( - status_code=status.HTTP_502_BAD_GATEWAY, - detail=f"TTS 合成失败: {e}", - ) from e - except ValueError as e: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail=str(e), - ) from e + except (CosyVoiceError, ValueError) as e: + # 合成失败退费 + if _points_deducted > 0 and _points_svc is not None: + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db) + except Exception as refund_err: + logger.warning(f"TTS 预览失败退积分异常: {refund_err}") + if isinstance(e, CosyVoiceError): + raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=f"TTS 合成失败: {e}") from e + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e return TTSPreviewResponse( audio_url=result.audio_url, diff --git a/apps/api/app/api/routes/voice_clones.py b/apps/api/app/api/routes/voice_clones.py index 82f8a0b44..25a35e041 100755 --- a/apps/api/app/api/routes/voice_clones.py +++ b/apps/api/app/api/routes/voice_clones.py @@ -3,14 +3,17 @@ from __future__ import annotations import logging +import math from typing import Optional from app.auth import AuthenticatedUser, get_current_user +from app.config import settings from app.core.celery_app import celery_app from app.core.storage import get_storage_service from app.dependencies import ( get_asset_repository, get_cosyvoice_service, + get_db_session, get_project_repository, get_voice_clone_profile_repository, ) @@ -22,6 +25,7 @@ from app.schemas.voice_clone import ( VoiceCloneStatusResponse, ) from fastapi import APIRouter, Depends, HTTPException, Query, Response, status +from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import ( SQLAlchemyVoiceCloneProfileRepository, @@ -38,6 +42,11 @@ from packages.application.voice_clone.use_cases import ( from packages.application.voice_clone.workflow import ( VoiceCloneWorkflowService, ) +from packages.domain.points_rules import calculate_points_cost +from packages.domain.points_service import PointsService + +# remove duplicate +_DUMMY_DELETED = () from packages.ports.asset_repository import AssetRepository from packages.ports.project_repository import ProjectRepository from packages.shared.storage import SharedStorageService @@ -339,6 +348,7 @@ def get_voice_clone_preview( description="情绪:neutral/happy/sad/angry/surprised/fearful/disgusted,兼容旧值 natural/excited/calm/friendly,空为默认自然", ), authenticated_user: AuthenticatedUser = Depends(get_current_user), + db: Session = Depends(get_db_session), repository: SQLAlchemyVoiceCloneProfileRepository = Depends(get_voice_clone_profile_repository), cosyvoice: CosyVoiceService = Depends(get_cosyvoice_service), ) -> VoiceClonePreviewResponse: @@ -350,6 +360,31 @@ def get_voice_clone_preview( """ import time + user_id = authenticated_user.user.id + _points_deducted = 0 + _points_scene = "voice_clone_synth" + _points_svc = PointsService() if settings.points_enabled else None + _preview_text_for_points = text.strip() or CLONE_PREVIEW_TEMPLATE + if _points_svc is not None: + est_minutes = max(1.0, math.ceil(len(_preview_text_for_points) / 240)) + _points_deducted = calculate_points_cost( + _points_scene, + is_member=getattr(authenticated_user.user, "is_member", False), + duration_minutes=est_minutes, + member_type=getattr(authenticated_user.user, "member_type", None), + ) + _deduct_res = _points_svc.deduct_points(user_id, _points_deducted, _points_scene, db) + if not _deduct_res["success"]: + raise HTTPException( + status_code=402, + detail={ + "code": "INSUFFICIENT_POINTS", + "message": f"积分不足,需要 {_points_deducted} 积分,当前余额 {_deduct_res['balance']}", + "required": _points_deducted, + "balance": _deduct_res["balance"], + }, + ) + if emotion not in _ALLOWED_PREVIEW_EMOTIONS: raise HTTPException( status_code=status.HTTP_400_BAD_REQUEST, @@ -393,9 +428,14 @@ def get_voice_clone_preview( speed=speed, emotion=emotion, ) - except CosyVoiceError as e: - raise HTTPException(status_code=502, detail=f"TTS 合成失败: {e}") from e - except ValueError as e: + except (CosyVoiceError, ValueError) as e: + if _points_deducted > 0 and _points_svc is not None: + try: + _points_svc.refund_points(user_id, _points_deducted, _points_scene, db) + except Exception as refund_err: + logger.warning(f"克隆音色试听失败退积分异常: clone_id={clone_id}, err={refund_err}") + if isinstance(e, CosyVoiceError): + raise HTTPException(status_code=502, detail=f"TTS 合成失败: {e}") from e raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e # 缓存(仅默认参数组合) diff --git a/tests/unit/test_ai_avatar_render_points.py b/tests/unit/test_ai_avatar_render_points.py new file mode 100644 index 000000000..6ae5eea8d --- /dev/null +++ b/tests/unit/test_ai_avatar_render_points.py @@ -0,0 +1,48 @@ +"""AI数字人渲染 积分扣点单元测试 (#1895 P2 step 2.6)""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +import packages.middleware.points_gate as _pg_module + + +@pytest.fixture(autouse=True) +def _enable(monkeypatch): + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True) + yield + + +class TestAiAvatarRenderPoints: + def test_ai_digital_human_per_unit(self): + from packages.domain.points_rules import calculate_points_cost + + cost = calculate_points_cost("ai_digital_human", is_member=False, duration_minutes=1) + assert cost >= 15 + + def test_decorator_attached(self): + from app.api.routes.ai_avatar_render import create_render_job + + assert hasattr(create_render_job, "__wrapped__"), "missing @points_gate" + + def test_insufficient_raises_402(self): + from app.api.routes.ai_avatar_render import create_render_job + from app.schemas.ai_avatar_render import CreateAiAvatarRenderRequest + from fastapi import HTTPException + + db = MagicMock() + cu = MagicMock() + cu.user.id = "u1" + cu.user.is_member = False + cu.user.member_type = None + svc = MagicMock() + body = CreateAiAvatarRenderRequest(lipsync_job_id="lip1") + with patch("packages.domain.points_service.PointsService") as MS: + msvc = MagicMock() + msvc.deduct_points.return_value = {"success": False, "balance": 0} + MS.return_value = msvc + with pytest.raises(HTTPException) as ei: + create_render_job(body=body, current_user=cu, svc=svc, db=db) + assert ei.value.status_code == 402 diff --git a/tests/unit/test_ai_avatar_render_routes.py b/tests/unit/test_ai_avatar_render_routes.py index 78f632b00..542fe622e 100644 --- a/tests/unit/test_ai_avatar_render_routes.py +++ b/tests/unit/test_ai_avatar_render_routes.py @@ -10,6 +10,16 @@ from unittest.mock import MagicMock, patch import pytest +import packages.middleware.points_gate as _pg_module + + +@pytest.fixture(autouse=True) +def _disable_points_gate(monkeypatch): + """默认关闭积分闸门,避免影响既有用例。""" + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: False) + yield + + os.environ.setdefault("JWT_SECRET_KEY", "dev-secret-key-for-testing") diff --git a/tests/unit/test_generation_preview.py b/tests/unit/test_generation_preview.py index 8da468eb5..69523f2bd 100644 --- a/tests/unit/test_generation_preview.py +++ b/tests/unit/test_generation_preview.py @@ -573,6 +573,15 @@ from app.schemas.generation_task import ( PreviewGenerationTaskResponse, ) +import packages.middleware.points_gate as _pg_module + + +@pytest.fixture(autouse=True) +def _disable_points_gate(monkeypatch): + """默认关闭积分闸门,避免影响既有用例。""" + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: False) + yield + def _make_user(user_id="test_user_001"): """构造 mock AuthenticatedUser""" diff --git a/tests/unit/test_generation_preview_points.py b/tests/unit/test_generation_preview_points.py new file mode 100644 index 000000000..196800d51 --- /dev/null +++ b/tests/unit/test_generation_preview_points.py @@ -0,0 +1,54 @@ +"""视频预览生成 积分扣点单元测试 (#1895 P2 step 2.5)""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +import packages.middleware.points_gate as _pg_module + + +@pytest.fixture(autouse=True) +def _enable(monkeypatch): + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True) + yield + + +class TestGenerationPreviewPoints: + def test_ai_video_cost(self): + from packages.domain.points_rules import calculate_points_cost + + assert calculate_points_cost("ai_video", is_member=False) == 4 + assert calculate_points_cost("ai_video", is_member=True, member_type="monthly") == 2 + + def test_insufficient_raises_402(self): + from app.api.routes.generation_preview import create_preview_generation_task + from app.schemas.generation_task import CreatePreviewGenerationTaskRequest + from fastapi import HTTPException + + db = MagicMock() + cu = MagicMock() + cu.user.id = "u1" + cu.user.is_member = False + cu.user.member_type = None + req = CreatePreviewGenerationTaskRequest(template_id="t1", asset_ids=["a1"], preview_count=1) + with patch("packages.domain.points_service.PointsService") as MS: + svc = MagicMock() + svc.check_daily_free_clip.return_value = False + svc.deduct_points.return_value = {"success": False, "balance": 0} + MS.return_value = svc + with pytest.raises(HTTPException) as ei: + create_preview_generation_task( + request=req, + authenticated_user=cu, + db=db, + generation_task_repository=MagicMock(), + asset_repo=MagicMock(), + ) + assert ei.value.status_code == 402 + + def test_decorator_attached(self): + from app.api.routes.generation_preview import create_preview_generation_task + + assert hasattr(create_preview_generation_task, "__wrapped__"), "missing @points_gate" diff --git a/tests/unit/test_generation_tasks.py b/tests/unit/test_generation_tasks.py index ad33cc4c3..239e1480d 100755 --- a/tests/unit/test_generation_tasks.py +++ b/tests/unit/test_generation_tasks.py @@ -6,6 +6,7 @@ from unittest.mock import MagicMock import pytest +import packages.middleware.points_gate as _pg_module from packages.application.generation_tasks import ( CreateGenerationTaskCommand, CreateGenerationTaskUseCase, @@ -18,6 +19,13 @@ from packages.application.generation_tasks import ( from packages.domain import GenerationTask +@pytest.fixture(autouse=True) +def _disable_points_gate(monkeypatch): + """默认关闭积分闸门,避免影响既有用例。""" + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: False) + yield + + @pytest.fixture def mock_repo(): return MagicMock() diff --git a/tests/unit/test_generation_tasks_points.py b/tests/unit/test_generation_tasks_points.py new file mode 100644 index 000000000..2599398ad --- /dev/null +++ b/tests/unit/test_generation_tasks_points.py @@ -0,0 +1,63 @@ +"""视频生成 积分扣点单元测试 (#1895 P2 step 2.4)""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest + +import packages.middleware.points_gate as _pg_module + + +@pytest.fixture(autouse=True) +def _enable(monkeypatch): + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True) + yield + + +class TestGenerationTasksPoints: + def test_ai_video_base_cost(self): + from packages.domain.points_rules import calculate_points_cost + + assert calculate_points_cost("ai_video", is_member=False) == 4 + assert calculate_points_cost("ai_video", is_member=True, member_type="monthly") == 2 + + def test_ai_video_quantity_scales(self): + from packages.domain.points_rules import calculate_points_cost + + c1 = calculate_points_cost("ai_video", is_member=False, quantity=1) + c3 = calculate_points_cost("ai_video", is_member=False, quantity=3) + assert c3 > c1 + + def test_insufficient_raises_402(self): + from app.api.routes.generation_tasks import create_generation_task + from app.schemas.generation_task import CreateGenerationTaskRequest + from fastapi import HTTPException + + db = MagicMock() + cu = MagicMock() + cu.user.id = "u1" + cu.user.is_member = False + cu.user.member_type = None + req = CreateGenerationTaskRequest(template_id="t1", asset_ids=["a1"], count=1) + with patch("packages.domain.points_service.PointsService") as MS: + svc = MagicMock() + svc.check_daily_free_clip.return_value = False + svc.deduct_points.return_value = {"success": False, "balance": 0} + MS.return_value = svc + with pytest.raises(HTTPException) as ei: + create_generation_task( + request=req, + authenticated_user=cu, + db=db, + generation_task_repository=MagicMock(), + project_repository=MagicMock(), + asset_library_repository=MagicMock(), + asset_repository=MagicMock(), + ) + assert ei.value.status_code == 402 + + def test_decorator_attached(self): + from app.api.routes.generation_tasks import create_generation_task + + assert hasattr(create_generation_task, "__wrapped__"), "missing @points_gate" diff --git a/tests/unit/test_lipsync_points.py b/tests/unit/test_lipsync_points.py new file mode 100644 index 000000000..5acb81789 --- /dev/null +++ b/tests/unit/test_lipsync_points.py @@ -0,0 +1,242 @@ +"""lipsync 积分扣点单元测试 (#1895 P2 step 2.2)""" + +from __future__ import annotations + +import math +from unittest.mock import MagicMock + +import pytest +from fastapi import HTTPException + + +def _make_cu(user_id="user-1", is_member=False, member_type=None): + cu = MagicMock() + cu.user.id = user_id + cu.user.is_member = is_member + cu.user.member_type = member_type + return cu + + +class TestLipsyncDurationEstimate: + @pytest.mark.parametrize( + "text,expected", + [ + ("你好", 1.0), + ("你" * 240, 1.0), + ("你" * 241, 2.0), + ("你" * 1000, 5.0), + ], + ) + def test_text_estimate(self, text, expected): + est = max(1.0, math.ceil(len(text) / 240)) + assert est == expected + + @pytest.mark.parametrize( + "seconds,expected", + [ + (30, 1.0), + (60, 1.0), + (61, 2.0), + (120, 2.0), + (180, 3.0), + ], + ) + def test_audio_duration_estimate(self, seconds, expected): + est = max(1.0, math.ceil(seconds / 60.0)) + assert est == expected + + +class TestLipsyncPointsDeduction: + def _deduct(self, text="你好", audio_duration=None, enabled=True, success=True, balance=100, **cu_kw): + from packages.domain.points_rules import calculate_points_cost + + svc = MagicMock() if enabled else None + cu = _make_cu(**cu_kw) + if svc is None: + return 0, cu + if audio_duration and audio_duration > 0: + est = max(1.0, math.ceil(audio_duration / 60.0)) + elif text: + est = max(1.0, math.ceil(len(text) / 240)) + else: + est = 1.0 + cost = calculate_points_cost( + "ai_digital_human", + is_member=getattr(cu.user, "is_member", False), + duration_minutes=est, + member_type=getattr(cu.user, "member_type", None), + ) + svc.deduct_points.return_value = {"success": success, "balance": balance} + res = svc.deduct_points(cu.user.id, cost, "ai_digital_human", MagicMock()) + if not res["success"]: + raise HTTPException(status_code=402, detail={"code": "INSUFFICIENT_POINTS"}) + return cost, cu + + def test_disabled(self): + cost, _ = self._deduct(enabled=False) + assert cost == 0 + + def test_short_text_min_1min(self): + cost, _ = self._deduct(text="你好") + assert cost >= 15 # 15 base/min for free user × 1.15 + + def test_audio_duration_used(self): + cost_long, _ = self._deduct(audio_duration=180) # 3min + cost_short, _ = self._deduct(audio_duration=30) # 1min + assert cost_long > cost_short + + def test_insufficient_402(self): + with pytest.raises(HTTPException) as ei: + self._deduct(text="你" * 500, success=False, balance=0) + assert ei.value.status_code == 402 + + def test_member_cheaper(self): + cm, _ = self._deduct(text="你" * 500, is_member=True, member_type="yearly") + cf, _ = self._deduct(text="你" * 500, is_member=False) + assert cm < cf + + +# ── 直接调用 create_lipsync_job 覆盖扣点/402/退费分支 ── +import importlib +from types import SimpleNamespace +from unittest.mock import patch + +import packages.middleware.points_gate as _pg_module + + +# Ensure the enable-gate fixture for lipsync also covers @points_gate (if any) +# (the existing autouse _enable is below; importlib to avoid duplicate) +def _do_enable(monkeypatch): + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True) + + +def _body(**kw): + b = MagicMock() + defaults = dict( + video_url="http://x/v.mp4", + audio_url=None, + audio_duration=None, + sentence_timings=None, + voice_id=None, + script_text="你好世界", + speed=1.0, + emotion="", + enable_video_loop=False, + project_id=None, + ) + defaults.update(kw) + for k, v in defaults.items(): + setattr(b, k, v) + return b + + +def _cu(user_id="u1", is_member=False, member_type=None): + cu = MagicMock() + cu.user.id = user_id + cu.user.is_member = is_member + cu.user.member_type = member_type + return cu + + +class TestLipsyncEndpointPoints: + def test_insufficient_raises_402(self, monkeypatch): + _do_enable(monkeypatch) + from app.api.routes.lipsync import create_lipsync_job + + db = MagicMock() + svc = MagicMock() + ps = MagicMock() + ps.deduct_points.return_value = {"success": False, "balance": 0} + fs = MagicMock(points_enabled=True) + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): + with pytest.raises(HTTPException) as ei: + create_lipsync_job(body=_body(script_text="你" * 500), current_user=_cu(), db=db, svc=svc) + assert ei.value.status_code == 402 + + def test_value_error_refunds(self, monkeypatch): + _do_enable(monkeypatch) + from app.api.routes.lipsync import create_lipsync_job + + db = MagicMock() + svc = MagicMock() + svc.create_job.side_effect = ValueError("bad input") + ps = MagicMock() + ps.deduct_points.return_value = {"success": True, "balance": 99} + fs = MagicMock(points_enabled=True) + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): + with pytest.raises(HTTPException) as ei: + create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc) + assert ei.value.status_code == 400 + assert ps.refund_points.called + + def test_mediakit_error_refunds(self, monkeypatch): + _do_enable(monkeypatch) + from app.api.routes.lipsync import create_lipsync_job + from app.services.mediakit_client import MediaKitError + + db = MagicMock() + svc = MagicMock() + svc.create_job.side_effect = MediaKitError("fail", code="InvalidInput") + ps = MagicMock() + ps.deduct_points.return_value = {"success": True, "balance": 99} + fs = MagicMock(points_enabled=True) + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): + with pytest.raises(HTTPException) as ei: + create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc) + assert ei.value.status_code == 400 + assert ps.refund_points.called + + def test_generic_exception_refunds(self, monkeypatch): + _do_enable(monkeypatch) + from app.api.routes.lipsync import create_lipsync_job + + db = MagicMock() + svc = MagicMock() + svc.create_job.side_effect = RuntimeError("boom") + ps = MagicMock() + ps.deduct_points.return_value = {"success": True, "balance": 99} + fs = MagicMock(points_enabled=True) + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): + with pytest.raises(HTTPException) as ei: + create_lipsync_job(body=_body(), current_user=_cu(), db=db, svc=svc) + assert ei.value.status_code == 400 + assert ps.refund_points.called + + def test_audio_duration_estimation(self, monkeypatch): + _do_enable(monkeypatch) + from app.api.routes.lipsync import create_lipsync_job + + from packages.domain.points_rules import calculate_points_cost + + db = MagicMock() + svc = MagicMock() + job = SimpleNamespace(id="job-1", status="queued") + svc.create_job.return_value = job + ps = MagicMock() + ps.deduct_points.return_value = {"success": True, "balance": 99} + fs = MagicMock(points_enabled=True) + with ( + patch("app.api.routes.lipsync.PointsService", return_value=ps), + patch("app.api.routes.lipsync.settings", fs), + ): + create_lipsync_job( + body=_body(audio_url="http://x/a.mp3", audio_duration=180, script_text=None), + current_user=_cu(), + db=db, + svc=svc, + ) + # 180 seconds -> 3 minutes; assert deduct called with cost >= 15*3 + args = ps.deduct_points.call_args[0] + assert args[1] >= calculate_points_cost("ai_digital_human", is_member=False, duration_minutes=3) diff --git a/tests/unit/test_scripts_ai.py b/tests/unit/test_scripts_ai.py index 3cf946e2b..87b9f2417 100644 --- a/tests/unit/test_scripts_ai.py +++ b/tests/unit/test_scripts_ai.py @@ -16,6 +16,8 @@ from unittest.mock import MagicMock, patch import pydantic import pytest +import packages.middleware.points_gate as _pg_module + sys.path.insert(0, "apps/api") @@ -28,6 +30,18 @@ def _make_auth_user(user_id: str = "u1"): return auth +@pytest.fixture(autouse=True) +def _disable_points_gate(monkeypatch): + """默认关闭积分闸门,避免影响既有用例。""" + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: False) + yield + + +@pytest.fixture +def mock_db(): + return MagicMock() + + def _mock_youtube_dl( extract_info_return=None, extract_info_side_effect=None, @@ -82,7 +96,7 @@ class TestExtractFromDouyin: req = ExtractFromDouyinRequest(url="https://v.douyin.com/xxxxx/") auth = _make_auth_user() - result = extract_from_douyin(request=req, authenticated_user=auth) + result = extract_from_douyin(request=req, current_user=auth) assert result.text == "这是一段测试文案内容" assert result.duration_seconds == 120.5 @@ -112,7 +126,7 @@ class TestExtractFromDouyin: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - extract_from_douyin(request=req, authenticated_user=auth) + extract_from_douyin(request=req, current_user=auth) assert exc_info.value.status_code == 400 @patch("tempfile.TemporaryDirectory") @@ -136,7 +150,7 @@ class TestExtractFromDouyin: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - extract_from_douyin(request=req, authenticated_user=auth) + extract_from_douyin(request=req, current_user=auth) assert exc_info.value.status_code == 502 @patch("app.api.routes.scripts_ai.transcribe_to_text") @@ -169,7 +183,7 @@ class TestExtractFromDouyin: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - extract_from_douyin(request=req, authenticated_user=auth) + extract_from_douyin(request=req, current_user=auth) assert exc_info.value.status_code == 503 @patch("app.api.routes.scripts_ai.transcribe_to_text") @@ -202,7 +216,7 @@ class TestExtractFromDouyin: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - extract_from_douyin(request=req, authenticated_user=auth) + extract_from_douyin(request=req, current_user=auth) assert exc_info.value.status_code == 502 @@ -225,7 +239,7 @@ class TestAiRewrite: req = AiRewriteRequest(content="原始文案内容", style="口语化") auth = _make_auth_user() - result = ai_rewrite(request=req, authenticated_user=auth) + result = ai_rewrite(request=req, current_user=auth) assert result.original == "原始文案内容" assert result.rewritten == "改写后的文案内容,口语化风格" @@ -242,7 +256,7 @@ class TestAiRewrite: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - ai_rewrite(request=req, authenticated_user=auth) + ai_rewrite(request=req, current_user=auth) assert exc_info.value.status_code == 400 @patch("app.api.routes.scripts_ai.get_doubao_client") @@ -261,7 +275,7 @@ class TestAiRewrite: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - ai_rewrite(request=req, authenticated_user=auth) + ai_rewrite(request=req, current_user=auth) assert exc_info.value.status_code == 502 @patch("app.api.routes.scripts_ai.get_doubao_client") @@ -279,7 +293,7 @@ class TestAiRewrite: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - ai_rewrite(request=req, authenticated_user=auth) + ai_rewrite(request=req, current_user=auth) assert exc_info.value.status_code == 502 @@ -301,7 +315,7 @@ class TestAiGenerateTitles: req = AiGenerateTitlesRequest(content="这是一段关于美食的文案", count=3) auth = _make_auth_user() - result = ai_generate_titles(request=req, authenticated_user=auth) + result = ai_generate_titles(request=req, current_user=auth) assert len(result.titles) == 3 assert all(isinstance(t, str) for t in result.titles) @@ -339,12 +353,12 @@ class TestAiGenerateTitles: # count=5 req = AiGenerateTitlesRequest(content="测试内容", count=5) - result = ai_generate_titles(request=req, authenticated_user=auth) + result = ai_generate_titles(request=req, current_user=auth) assert len(result.titles) <= 5 # count=1 req = AiGenerateTitlesRequest(content="测试内容", count=1) - result = ai_generate_titles(request=req, authenticated_user=auth) + result = ai_generate_titles(request=req, current_user=auth) assert len(result.titles) >= 1 def test_generate_titles_empty_content(self): @@ -357,7 +371,7 @@ class TestAiGenerateTitles: auth = _make_auth_user() with pytest.raises(HTTPException) as exc_info: - ai_generate_titles(request=req, authenticated_user=auth) + ai_generate_titles(request=req, current_user=auth) assert exc_info.value.status_code == 400 @patch("app.services.ai_service.get_doubao_client") @@ -372,7 +386,7 @@ class TestAiGenerateTitles: req = AiGenerateTitlesRequest(content="测试文案内容") auth = _make_auth_user() - result = ai_generate_titles(request=req, authenticated_user=auth) + result = ai_generate_titles(request=req, current_user=auth) assert len(result.titles) == 3 diff --git a/tests/unit/test_scripts_ai_points.py b/tests/unit/test_scripts_ai_points.py new file mode 100644 index 000000000..70967fafa --- /dev/null +++ b/tests/unit/test_scripts_ai_points.py @@ -0,0 +1,76 @@ +"""scripts_ai 积分扣点单元测试 (#1895 P2 step 2.3)""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import pytest +from fastapi import HTTPException + +import packages.middleware.points_gate as _pg_module + + +def _make_cu(user_id="u1", is_member=False, member_type=None): + cu = MagicMock() + cu.user.id = user_id + cu.user.is_member = is_member + cu.user.member_type = member_type + return cu + + +@pytest.fixture(autouse=True) +def _enable_gate(monkeypatch): + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True) + yield + + +class TestScriptsAiPointsGate: + """测试 scripts_ai 三个端点都挂了 @points_gate 并正确扣费。""" + + @pytest.mark.parametrize( + "scene,endpoint_fn_name", + [ + ("douyin_extract", "extract_from_douyin"), + ("ai_rewrite", "ai_rewrite"), + ("ai_title", "ai_generate_titles"), + ], + ) + def test_insufficient_points_raises_402(self, scene, endpoint_fn_name): + """积分不足时抛 402。""" + from app.api.routes import scripts_ai + from app.schemas.scripts_ai import ( + AiGenerateTitlesRequest, + AiRewriteRequest, + ExtractFromDouyinRequest, + ) + + fn = getattr(scripts_ai, endpoint_fn_name) + db = MagicMock() + cu = _make_cu() + if scene == "douyin_extract": + req = ExtractFromDouyinRequest(url="https://v.douyin.com/abc/") + elif scene == "ai_rewrite": + req = AiRewriteRequest(content="测试文案") + else: + req = AiGenerateTitlesRequest(content="测试文案", count=3) + + with patch("packages.domain.points_service.PointsService") as MockSvc: + svc = MagicMock() + svc.deduct_points.return_value = {"success": False, "balance": 0} + MockSvc.return_value = svc + with pytest.raises(HTTPException) as ei: + fn(request=req, current_user=cu, db=db) + assert ei.value.status_code == 402 + + def test_disabled_passthrough_no_user_error(self, monkeypatch): + """关闭时不需要 user/db 也能被装饰器透传(验证 gate 关闭零副作用)。""" + from app.api.routes import scripts_ai + from app.schemas.scripts_ai import AiRewriteRequest + + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: False) + fn = scripts_ai.ai_rewrite + # 不带 db/current_user 也应透传(后续业务逻辑可能报错但不是 401/500 gate 错误) + with pytest.raises(Exception) as ei: + fn(request=AiRewriteRequest(content="x"), current_user=None, db=None) + # 不应是 gate 抛的 401/500 + assert isinstance(ei.value, AttributeError) or ei.value.status_code not in (401, 500) diff --git a/tests/unit/test_tts_voice_clone_points.py b/tests/unit/test_tts_voice_clone_points.py new file mode 100644 index 000000000..8e5b65c59 --- /dev/null +++ b/tests/unit/test_tts_voice_clone_points.py @@ -0,0 +1,267 @@ +"""TTS + voice_clone 积分扣点单元测试 (#1895 P2 step 2.1) + +覆盖 synthesize / voice_clone preview 在积分开关下的扣点、余额不足、失败退费、会员折扣等分支。 +""" + +from __future__ import annotations + +import math +from types import SimpleNamespace +from unittest.mock import MagicMock, patch + +import pytest +from fastapi import HTTPException + +import packages.middleware.points_gate as _pg_module + + +@pytest.fixture(autouse=True) +def _enable_gate(monkeypatch): + monkeypatch.setattr(_pg_module, "_points_gate_enabled", lambda: True) + yield + + +def _make_cu(user_id="user-1", is_member=False, member_type=None): + cu = MagicMock() + cu.user.id = user_id + cu.user.is_member = is_member + cu.user.member_type = member_type + return cu + + +def _make_request(text="你好世界", voice_id="v1", **kw): + r = MagicMock() + r.text = text + r.voice_id = voice_id + r.voice_clone_profile_id = None + r.speed = 1.0 + r.emotion = "" + r.language = "zh-CN" + r.metadata_ = {} + r.voice_model = None + for k, v in kw.items(): + setattr(r, k, v) + return r + + +def _est_minutes(chars: int) -> float: + return max(1.0, math.ceil(chars / 240)) + + +class TestEstimateMinutes: + @pytest.mark.parametrize( + "chars,expected", + [(1, 1.0), (240, 1.0), (241, 2.0), (480, 2.0), (481, 3.0), (1000, 5.0)], + ) + def test_estimate(self, chars, expected): + assert _est_minutes(chars) == expected + + +class TestTtsSynthesizePointsDeduction: + def _setup(self, text="你好", deduct_success=True, balance=0, start_synth_raises=None, send_task_raises=None): + db = MagicMock() + cu = _make_cu() + repo = MagicMock() + import enum + + class _S(enum.Enum): + processing = "processing" + + job = SimpleNamespace(id="job-1", status=_S.processing, metadata={}) + uc = MagicMock() + uc.execute.return_value = job + wf = MagicMock() + wf.start_synthesis.return_value = job + wf.process_synthesis_failure.return_value = job + if start_synth_raises: + wf.start_synthesis.side_effect = start_synth_raises + vc_repo = MagicMock() + vc_repo.get.return_value = None + svc = MagicMock() + svc.deduct_points.return_value = {"success": deduct_success, "balance": balance} + fake_settings = MagicMock(points_enabled=True) + return db, cu, repo, uc, wf, vc_repo, svc, fake_settings, job + + def test_insufficient_raises_402(self): + db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup(text="你好" * 200, deduct_success=False, balance=0) + from app.api.routes.tts import synthesize + + with ( + patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), + patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), + patch("app.api.routes.tts.PointsService", return_value=svc), + patch("app.api.routes.tts.settings", fs), + ): + with pytest.raises(HTTPException) as ei: + synthesize( + request=_make_request(text="你好" * 200), + authenticated_user=cu, + db=db, + repository=repo, + cosyvoice_service=MagicMock(), + voice_clone_repo=vc_repo, + ) + assert ei.value.status_code == 402 + assert ei.value.detail["code"] == "INSUFFICIENT_POINTS" + + def test_success_deducts_points(self): + db, cu, repo, uc, wf, vc_repo, svc, fs, job = self._setup() + from app.api.routes.tts import synthesize + + with ( + patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), + patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), + patch("app.api.routes.tts.PointsService", return_value=svc), + patch("app.api.routes.tts.celery_app.send_task") as _st, + patch("app.api.routes.tts.settings", fs), + ): + resp = synthesize( + request=_make_request(text="测试"), + authenticated_user=cu, + db=db, + repository=repo, + cosyvoice_service=MagicMock(), + voice_clone_repo=vc_repo, + ) + svc.deduct_points.assert_called_once() + assert resp.job_id == job.id + + def test_synthesis_failure_refunds(self): + db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup(start_synth_raises=RuntimeError("boom")) + from app.api.routes.tts import synthesize + + with ( + patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), + patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), + patch("app.api.routes.tts.PointsService", return_value=svc), + patch("app.api.routes.tts.celery_app.send_task"), + patch("app.api.routes.tts.settings", fs), + ): + synthesize( + request=_make_request(text="测试"), + authenticated_user=cu, + db=db, + repository=repo, + cosyvoice_service=MagicMock(), + voice_clone_repo=vc_repo, + ) + assert svc.refund_points.called + + def test_celery_send_failure_refunds(self): + db, cu, repo, uc, wf, vc_repo, svc, fs, _ = self._setup(send_task_raises=RuntimeError("celery down")) + from app.api.routes.tts import synthesize + + with ( + patch("app.api.routes.tts.CreateTTSJobUseCase", return_value=uc), + patch("app.api.routes.tts.TTSWorkflowService", return_value=wf), + patch("app.api.routes.tts.PointsService", return_value=svc), + patch("app.api.routes.tts.celery_app.send_task", side_effect=RuntimeError("celery down")), + patch("app.api.routes.tts.settings", fs), + ): + synthesize( + request=_make_request(text="测试"), + authenticated_user=cu, + db=db, + repository=repo, + cosyvoice_service=MagicMock(), + voice_clone_repo=vc_repo, + ) + assert svc.refund_points.called + + def test_member_cheaper(self): + from packages.domain.points_rules import calculate_points_cost + + cf = calculate_points_cost("ai_voice", is_member=False, duration_minutes=2) + cm = calculate_points_cost("ai_voice", is_member=True, member_type="monthly", duration_minutes=2) + assert cm < cf + + +class TestVoiceClonePreviewPoints: + def _setup(self, text="你好", deduct_success=True, balance=0, synth_raises=None): + db = MagicMock() + cu = _make_cu() + repo = MagicMock() + profile = SimpleNamespace(is_ready=True, voice_id="vc-1", user_id=cu.user.id) + uc = MagicMock() + uc.execute.return_value = profile + cosy = MagicMock() + r = SimpleNamespace(audio_url="http://x/a.mp3", duration=1.2, file_size=1000) + cosy.synthesize_speech.return_value = r + if synth_raises: + cosy.synthesize_speech.side_effect = synth_raises + svc = MagicMock() + svc.deduct_points.return_value = {"success": deduct_success, "balance": balance} + fs = MagicMock(points_enabled=True) + return db, cu, repo, uc, cosy, svc, fs + + def test_insufficient_raises_402(self): + db, cu, repo, uc, cosy, svc, fs = self._setup(deduct_success=False, balance=0) + from app.api.routes.voice_clones import get_voice_clone_preview + + with ( + patch("app.api.routes.voice_clones.GetVoiceCloneUseCase", return_value=uc), + patch("app.api.routes.voice_clones.PointsService", return_value=svc), + patch("app.api.routes.voice_clones.settings", fs), + patch("app.api.routes.voice_clones._clone_preview_cache", {}), + ): + with pytest.raises(HTTPException) as ei: + get_voice_clone_preview( + clone_id="c1", + text="你好", + speed=1.0, + emotion="", + authenticated_user=cu, + db=db, + repository=repo, + cosyvoice=cosy, + ) + assert ei.value.status_code == 402 + + def test_synth_cosyvoice_error_refunds_and_raises_502(self): + from packages.application.cosyvoice_service import CosyVoiceError + + db, cu, repo, uc, cosy, svc, fs = self._setup(synth_raises=CosyVoiceError("fail")) + from app.api.routes.voice_clones import get_voice_clone_preview + + with ( + patch("app.api.routes.voice_clones.GetVoiceCloneUseCase", return_value=uc), + patch("app.api.routes.voice_clones.PointsService", return_value=svc), + patch("app.api.routes.voice_clones.settings", fs), + patch("app.api.routes.voice_clones._clone_preview_cache", {}), + ): + with pytest.raises(HTTPException) as ei: + get_voice_clone_preview( + clone_id="c1", + text="你好", + speed=1.0, + emotion="", + authenticated_user=cu, + db=db, + repository=repo, + cosyvoice=cosy, + ) + assert ei.value.status_code == 502 + assert svc.refund_points.called + + def test_success_returns_audio(self): + db, cu, repo, uc, cosy, svc, fs = self._setup() + from app.api.routes.voice_clones import get_voice_clone_preview + + with ( + patch("app.api.routes.voice_clones.GetVoiceCloneUseCase", return_value=uc), + patch("app.api.routes.voice_clones.PointsService", return_value=svc), + patch("app.api.routes.voice_clones.settings", fs), + patch("app.api.routes.voice_clones._clone_preview_cache", {}), + ): + resp = get_voice_clone_preview( + clone_id="c1", + text="你好", + speed=1.0, + emotion="", + authenticated_user=cu, + db=db, + repository=repo, + cosyvoice=cosy, + ) + svc.deduct_points.assert_called_once() + assert resp.audio_url.startswith("http")