Compare commits
12 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| 4607977496 | |||
| be1451dc21 | |||
| be9031abc9 | |||
| a96aff1ac4 | |||
| a7317bb8e4 | |||
| b5a593feb9 | |||
| b3ff3c9936 | |||
| 1cf9bbd9a5 | |||
| f2a13ef5a7 | |||
| 6084b2f662 | |||
| 3862d9158f | |||
| 580ec51928 |
@@ -29,6 +29,8 @@ from app.services.ai_avatar_render_service import (
|
|||||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||||
from sqlalchemy.orm import Session
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
|
from packages.middleware.points_gate import points_gate
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
@@ -42,10 +44,12 @@ def _get_service(db: Session = Depends(get_db_session)) -> AiAvatarRenderService
|
|||||||
|
|
||||||
|
|
||||||
@router.post("", response_model=AiAvatarRenderJobResponse, status_code=201)
|
@router.post("", response_model=AiAvatarRenderJobResponse, status_code=201)
|
||||||
|
@points_gate("ai_digital_human", per_unit=15)
|
||||||
def create_render_job(
|
def create_render_job(
|
||||||
body: CreateAiAvatarRenderRequest,
|
body: CreateAiAvatarRenderRequest,
|
||||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
svc: AiAvatarRenderService = Depends(_get_service),
|
svc: AiAvatarRenderService = Depends(_get_service),
|
||||||
|
db: Session = Depends(get_db_session),
|
||||||
):
|
):
|
||||||
"""提交 AI 数字人渲染任务.
|
"""提交 AI 数字人渲染任务.
|
||||||
|
|
||||||
|
|||||||
@@ -27,6 +27,7 @@ from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
|||||||
)
|
)
|
||||||
from packages.application import ListGeneratedVideosByTaskUseCase
|
from packages.application import ListGeneratedVideosByTaskUseCase
|
||||||
from packages.domain.config_schemas import normalize_plan_config
|
from packages.domain.config_schemas import normalize_plan_config
|
||||||
|
from packages.middleware.points_gate import points_gate
|
||||||
from packages.shared.storage import get_shared_storage_service
|
from packages.shared.storage import get_shared_storage_service
|
||||||
|
|
||||||
from .templates_editor.dependencies import get_draft_plan_id, get_editor_services
|
from .templates_editor.dependencies import get_draft_plan_id, get_editor_services
|
||||||
@@ -331,6 +332,7 @@ def _is_trusted_media_url(url: str) -> bool:
|
|||||||
|
|
||||||
|
|
||||||
@router.post("/generate-cover", response_model=GenerateCoverResponse)
|
@router.post("/generate-cover", response_model=GenerateCoverResponse)
|
||||||
|
@points_gate("ai_cover")
|
||||||
def generate_cover(
|
def generate_cover(
|
||||||
body: GenerateCoverRequest,
|
body: GenerateCoverRequest,
|
||||||
template_id: str = Query(..., description="模板 ID"),
|
template_id: str = Query(..., description="模板 ID"),
|
||||||
|
|||||||
@@ -43,6 +43,7 @@ from packages.application import (
|
|||||||
GetGenerationTaskUseCase,
|
GetGenerationTaskUseCase,
|
||||||
ListGeneratedVideosByTaskUseCase,
|
ListGeneratedVideosByTaskUseCase,
|
||||||
)
|
)
|
||||||
|
from packages.middleware.points_gate import points_gate
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
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)
|
@router.post("/preview", response_model=BatchPreviewGenerationTaskResponse, status_code=201)
|
||||||
|
@points_gate("ai_video", quantity_field="preview_count")
|
||||||
def create_preview_generation_task(
|
def create_preview_generation_task(
|
||||||
request: CreatePreviewGenerationTaskRequest,
|
request: CreatePreviewGenerationTaskRequest,
|
||||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
|||||||
@@ -42,6 +42,7 @@ from packages.application import (
|
|||||||
ListGeneratedVideosByTaskUseCase,
|
ListGeneratedVideosByTaskUseCase,
|
||||||
)
|
)
|
||||||
from packages.domain.smart_match import smart_select_assets
|
from packages.domain.smart_match import smart_select_assets
|
||||||
|
from packages.middleware.points_gate import points_gate
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -211,6 +212,7 @@ def _resolve_project_and_library(
|
|||||||
|
|
||||||
|
|
||||||
@router.post("/tasks", response_model=BatchGenerationTaskResponse)
|
@router.post("/tasks", response_model=BatchGenerationTaskResponse)
|
||||||
|
@points_gate("ai_video", quantity_field="count")
|
||||||
def create_generation_task(
|
def create_generation_task(
|
||||||
request: CreateGenerationTaskRequest,
|
request: CreateGenerationTaskRequest,
|
||||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
|||||||
@@ -12,9 +12,11 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import math
|
||||||
from datetime import UTC
|
from datetime import UTC
|
||||||
|
|
||||||
from app.auth import AuthenticatedUser, get_current_user
|
from app.auth import AuthenticatedUser, get_current_user
|
||||||
|
from app.config import settings
|
||||||
from app.dependencies import (
|
from app.dependencies import (
|
||||||
get_db_session,
|
get_db_session,
|
||||||
get_voice_clone_profile_repository,
|
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 fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query
|
||||||
from sqlalchemy.orm import Session
|
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__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
@@ -53,8 +58,40 @@ def _get_service(
|
|||||||
def create_lipsync_job(
|
def create_lipsync_job(
|
||||||
body: CreateLipsyncJobRequest,
|
body: CreateLipsyncJobRequest,
|
||||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
db: Session = Depends(get_db_session),
|
||||||
svc: LipsyncService = Depends(_get_service),
|
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:
|
try:
|
||||||
job = svc.create_job(
|
job = svc.create_job(
|
||||||
user_id=current_user.user.id,
|
user_id=user_id,
|
||||||
video_url=body.video_url,
|
video_url=body.video_url,
|
||||||
audio_url=body.audio_url,
|
audio_url=body.audio_url,
|
||||||
audio_duration=body.audio_duration,
|
audio_duration=body.audio_duration,
|
||||||
@@ -79,8 +116,18 @@ def create_lipsync_job(
|
|||||||
project_id=body.project_id,
|
project_id=body.project_id,
|
||||||
)
|
)
|
||||||
except ValueError as exc:
|
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
|
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||||
except MediaKitError as 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
|
status_code = 502
|
||||||
if exc.code in ("VoiceForbidden",):
|
if exc.code in ("VoiceForbidden",):
|
||||||
status_code = 403
|
status_code = 403
|
||||||
@@ -96,11 +143,24 @@ def create_lipsync_job(
|
|||||||
) from exc
|
) from exc
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error("创建对口型任务异常: %s", exc, exc_info=True)
|
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(
|
raise HTTPException(
|
||||||
status_code=400,
|
status_code=400,
|
||||||
detail=f"创建对口型任务失败: {exc}",
|
detail=f"创建对口型任务失败: {exc}",
|
||||||
) from 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
|
return job
|
||||||
|
|
||||||
|
|
||||||
@@ -111,8 +171,34 @@ def create_lipsync_job(
|
|||||||
def preview_tts(
|
def preview_tts(
|
||||||
body: AiAvatarTtsPreviewRequest,
|
body: AiAvatarTtsPreviewRequest,
|
||||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
db: Session = Depends(get_db_session),
|
||||||
svc: LipsyncService = Depends(_get_service),
|
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 预合成.
|
"""步骤1「生成配音」同步 TTS 预合成.
|
||||||
|
|
||||||
同步执行 TTS 合成 → 下载音频 → ffprobe 时长 → 句子时间戳计算,
|
同步执行 TTS 合成 → 下载音频 → ffprobe 时长 → 句子时间戳计算,
|
||||||
@@ -121,13 +207,18 @@ def preview_tts(
|
|||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
result = svc.preview_tts(
|
result = svc.preview_tts(
|
||||||
user_id=current_user.user.id,
|
user_id=user_id,
|
||||||
voice_id=body.voice_id,
|
voice_id=body.voice_id,
|
||||||
script_text=body.script_text,
|
script_text=body.script_text,
|
||||||
speed=body.speed,
|
speed=body.speed,
|
||||||
emotion=body.emotion,
|
emotion=body.emotion,
|
||||||
)
|
)
|
||||||
except MediaKitError as 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"TTS 预合成 MediaKitError 退积分异常: err={refund_err}")
|
||||||
status_code = 400
|
status_code = 400
|
||||||
if exc.code in ("VoiceForbidden",):
|
if exc.code in ("VoiceForbidden",):
|
||||||
status_code = 403
|
status_code = 403
|
||||||
@@ -142,6 +233,11 @@ def preview_tts(
|
|||||||
) from exc
|
) from exc
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error("TTS 预合成异常: %s", exc, exc_info=True)
|
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(
|
raise HTTPException(
|
||||||
status_code=400,
|
status_code=400,
|
||||||
detail=f"TTS 合成失败: {exc}",
|
detail=f"TTS 合成失败: {exc}",
|
||||||
|
|||||||
@@ -13,6 +13,7 @@ import re
|
|||||||
import tempfile
|
import tempfile
|
||||||
|
|
||||||
from app.auth import AuthenticatedUser, get_current_user
|
from app.auth import AuthenticatedUser, get_current_user
|
||||||
|
from app.dependencies import get_db_session
|
||||||
from app.schemas.scripts_ai import (
|
from app.schemas.scripts_ai import (
|
||||||
AiGenerateTitlesRequest,
|
AiGenerateTitlesRequest,
|
||||||
AiGenerateTitlesResponse,
|
AiGenerateTitlesResponse,
|
||||||
@@ -27,7 +28,9 @@ from app.services.script_asr_service import (
|
|||||||
transcribe_to_text,
|
transcribe_to_text,
|
||||||
)
|
)
|
||||||
from fastapi import APIRouter, Depends, HTTPException, status
|
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
|
from packages.shared.ai_client import get_doubao_client
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -62,9 +65,11 @@ def _validate_douyin_url(url: str) -> None:
|
|||||||
"/extract-from-douyin",
|
"/extract-from-douyin",
|
||||||
response_model=ExtractFromDouyinResponse,
|
response_model=ExtractFromDouyinResponse,
|
||||||
)
|
)
|
||||||
|
@points_gate("douyin_extract")
|
||||||
def extract_from_douyin(
|
def extract_from_douyin(
|
||||||
request: ExtractFromDouyinRequest,
|
request: ExtractFromDouyinRequest,
|
||||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
db: Session = Depends(get_db_session),
|
||||||
) -> ExtractFromDouyinResponse:
|
) -> ExtractFromDouyinResponse:
|
||||||
"""从抖音视频下载无水印视频并通过 ASR 提取文案."""
|
"""从抖音视频下载无水印视频并通过 ASR 提取文案."""
|
||||||
source_url = request.url.strip()
|
source_url = request.url.strip()
|
||||||
@@ -138,9 +143,11 @@ def extract_from_douyin(
|
|||||||
"/ai-rewrite",
|
"/ai-rewrite",
|
||||||
response_model=AiRewriteResponse,
|
response_model=AiRewriteResponse,
|
||||||
)
|
)
|
||||||
|
@points_gate("ai_rewrite")
|
||||||
def ai_rewrite(
|
def ai_rewrite(
|
||||||
request: AiRewriteRequest,
|
request: AiRewriteRequest,
|
||||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
db: Session = Depends(get_db_session),
|
||||||
) -> AiRewriteResponse:
|
) -> AiRewriteResponse:
|
||||||
"""使用豆包大模型改写文案."""
|
"""使用豆包大模型改写文案."""
|
||||||
content = (request.content or "").strip()
|
content = (request.content or "").strip()
|
||||||
@@ -206,9 +213,11 @@ def ai_rewrite(
|
|||||||
"/ai-generate-titles",
|
"/ai-generate-titles",
|
||||||
response_model=AiGenerateTitlesResponse,
|
response_model=AiGenerateTitlesResponse,
|
||||||
)
|
)
|
||||||
|
@points_gate("ai_title")
|
||||||
def ai_generate_titles(
|
def ai_generate_titles(
|
||||||
request: AiGenerateTitlesRequest,
|
request: AiGenerateTitlesRequest,
|
||||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
db: Session = Depends(get_db_session),
|
||||||
) -> AiGenerateTitlesResponse:
|
) -> AiGenerateTitlesResponse:
|
||||||
"""使用现有 generate_smart_titles 生成标题."""
|
"""使用现有 generate_smart_titles 生成标题."""
|
||||||
content = (request.content or "").strip()
|
content = (request.content or "").strip()
|
||||||
|
|||||||
@@ -4,12 +4,14 @@ from __future__ import annotations
|
|||||||
|
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
|
import math
|
||||||
import subprocess
|
import subprocess
|
||||||
import tempfile
|
import tempfile
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional
|
||||||
|
|
||||||
from app.auth import AuthenticatedUser, get_current_user
|
from app.auth import AuthenticatedUser, get_current_user
|
||||||
|
from app.config import settings
|
||||||
from app.core.celery_app import celery_app
|
from app.core.celery_app import celery_app
|
||||||
from app.core.storage import get_storage_service
|
from app.core.storage import get_storage_service
|
||||||
from app.dependencies import (
|
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.application.tts_job.workflow import TTSWorkflowService
|
||||||
from packages.domain import Asset, AssetLibrary, AssetLibraryKind, AssetStatus, ClassificationStatus
|
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.domain.voice_presets import list_voices
|
||||||
from packages.ports.asset_library_repository import AssetLibraryRepository
|
from packages.ports.asset_library_repository import AssetLibraryRepository
|
||||||
from packages.ports.asset_repository import AssetRepository
|
from packages.ports.asset_repository import AssetRepository
|
||||||
@@ -128,6 +132,7 @@ def _to_response(job, sign_url=None) -> TTSJobResponse:
|
|||||||
def synthesize(
|
def synthesize(
|
||||||
request: TTSSynthesizeRequest,
|
request: TTSSynthesizeRequest,
|
||||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
db: Session = Depends(get_db_session),
|
||||||
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
|
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
|
||||||
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
|
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
|
||||||
voice_clone_repo=Depends(get_voice_clone_profile_repository),
|
voice_clone_repo=Depends(get_voice_clone_profile_repository),
|
||||||
@@ -139,6 +144,31 @@ def synthesize(
|
|||||||
"""
|
"""
|
||||||
user_id = authenticated_user.user.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:
|
||||||
|
# 中文按 ~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),
|
# 解析 voice_id:前端可能传克隆音色 profile UUID(而非 CosyVoice voice_id),
|
||||||
# 与 /tts/preview 保持一致:命中 profile → 校验归属 → 取 CosyVoice voice_id
|
# 与 /tts/preview 保持一致:命中 profile → 校验归属 → 取 CosyVoice voice_id
|
||||||
actual_voice_id = request.voice_id
|
actual_voice_id = request.voice_id
|
||||||
@@ -198,6 +228,7 @@ def synthesize(
|
|||||||
cosyvoice_service=cosyvoice_service,
|
cosyvoice_service=cosyvoice_service,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
synthesis_error: Exception | None = None
|
||||||
try:
|
try:
|
||||||
job = workflow.start_synthesis(job.id)
|
job = workflow.start_synthesis(job.id)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
@@ -205,10 +236,17 @@ def synthesize(
|
|||||||
# 但 DB 异常、网络异常等意外错误可能逃逸。
|
# 但 DB 异常、网络异常等意外错误可能逃逸。
|
||||||
# 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。
|
# 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。
|
||||||
logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True)
|
logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True)
|
||||||
|
synthesis_error = e
|
||||||
try:
|
try:
|
||||||
job = workflow.process_synthesis_failure(job.id, str(e))
|
job = workflow.process_synthesis_failure(job.id, str(e))
|
||||||
except Exception as inner_e:
|
except Exception as inner_e:
|
||||||
logger.error(f"标记 TTS job 失败时出错: job_id={job.id}, error={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 后台轮询
|
# 若任务处于 processing 状态(异步模式),触发 Celery 后台轮询
|
||||||
if job.status.value == "processing":
|
if job.status.value == "processing":
|
||||||
@@ -223,10 +261,17 @@ def synthesize(
|
|||||||
celery_app.send_task("worker.process_tts_synthesis", args=[job.id])
|
celery_app.send_task("worker.process_tts_synthesis", args=[job.id])
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
# Celery 调度失败,标记 job 为 failed
|
# Celery 调度失败,标记 job 为 failed
|
||||||
|
# e used below for refund context
|
||||||
try:
|
try:
|
||||||
workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}")
|
workflow.process_synthesis_failure(job.id, f"Celery 任务调度失败: {e}")
|
||||||
except Exception as inner_e:
|
except Exception as inner_e:
|
||||||
logger.error(f"Celery 调度后标记失败时出错: job_id={job.id}, error={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(
|
return TTSSynthesizeResponse(
|
||||||
job_id=job.id,
|
job_id=job.id,
|
||||||
@@ -553,6 +598,7 @@ def save_tts_job_to_library(
|
|||||||
def preview_tts(
|
def preview_tts(
|
||||||
request: TTSPreviewRequest,
|
request: TTSPreviewRequest,
|
||||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
db: Session = Depends(get_db_session),
|
||||||
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
|
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
|
||||||
voice_clone_repo=Depends(get_voice_clone_profile_repository),
|
voice_clone_repo=Depends(get_voice_clone_profile_repository),
|
||||||
) -> TTSPreviewResponse:
|
) -> TTSPreviewResponse:
|
||||||
@@ -561,6 +607,31 @@ def preview_tts(
|
|||||||
用于前端预览配音效果,限制文本长度 200 字以内。
|
用于前端预览配音效果,限制文本长度 200 字以内。
|
||||||
支持预设音色和克隆音色:克隆音色传的是 profile UUID,需解析为 CosyVoice voice_id。
|
支持预设音色和克隆音色:克隆音色传的是 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
|
# 解析 voice_id:前端可能传 VoiceCloneProfile UUID 或预设音色 ID
|
||||||
actual_voice_id = request.voice_id
|
actual_voice_id = request.voice_id
|
||||||
profile = voice_clone_repo.get(request.voice_id)
|
profile = voice_clone_repo.get(request.voice_id)
|
||||||
@@ -586,16 +657,16 @@ def preview_tts(
|
|||||||
emotion=request.emotion,
|
emotion=request.emotion,
|
||||||
language=getattr(request, "language", "zh-CN"),
|
language=getattr(request, "language", "zh-CN"),
|
||||||
)
|
)
|
||||||
except CosyVoiceError as e:
|
except (CosyVoiceError, ValueError) as e:
|
||||||
raise HTTPException(
|
# 合成失败退费
|
||||||
status_code=status.HTTP_502_BAD_GATEWAY,
|
if _points_deducted > 0 and _points_svc is not None:
|
||||||
detail=f"TTS 合成失败: {e}",
|
try:
|
||||||
) from e
|
_points_svc.refund_points(user_id, _points_deducted, _points_scene, db)
|
||||||
except ValueError as e:
|
except Exception as refund_err:
|
||||||
raise HTTPException(
|
logger.warning(f"TTS 预览失败退积分异常: {refund_err}")
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
if isinstance(e, CosyVoiceError):
|
||||||
detail=str(e),
|
raise HTTPException(status_code=status.HTTP_502_BAD_GATEWAY, detail=f"TTS 合成失败: {e}") from e
|
||||||
) from e
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
|
||||||
|
|
||||||
return TTSPreviewResponse(
|
return TTSPreviewResponse(
|
||||||
audio_url=result.audio_url,
|
audio_url=result.audio_url,
|
||||||
|
|||||||
@@ -3,14 +3,17 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import math
|
||||||
from typing import Optional
|
from typing import Optional
|
||||||
|
|
||||||
from app.auth import AuthenticatedUser, get_current_user
|
from app.auth import AuthenticatedUser, get_current_user
|
||||||
|
from app.config import settings
|
||||||
from app.core.celery_app import celery_app
|
from app.core.celery_app import celery_app
|
||||||
from app.core.storage import get_storage_service
|
from app.core.storage import get_storage_service
|
||||||
from app.dependencies import (
|
from app.dependencies import (
|
||||||
get_asset_repository,
|
get_asset_repository,
|
||||||
get_cosyvoice_service,
|
get_cosyvoice_service,
|
||||||
|
get_db_session,
|
||||||
get_project_repository,
|
get_project_repository,
|
||||||
get_voice_clone_profile_repository,
|
get_voice_clone_profile_repository,
|
||||||
)
|
)
|
||||||
@@ -22,6 +25,7 @@ from app.schemas.voice_clone import (
|
|||||||
VoiceCloneStatusResponse,
|
VoiceCloneStatusResponse,
|
||||||
)
|
)
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
|
||||||
|
from sqlalchemy.orm import Session
|
||||||
|
|
||||||
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
|
from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
|
||||||
SQLAlchemyVoiceCloneProfileRepository,
|
SQLAlchemyVoiceCloneProfileRepository,
|
||||||
@@ -38,6 +42,11 @@ from packages.application.voice_clone.use_cases import (
|
|||||||
from packages.application.voice_clone.workflow import (
|
from packages.application.voice_clone.workflow import (
|
||||||
VoiceCloneWorkflowService,
|
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.asset_repository import AssetRepository
|
||||||
from packages.ports.project_repository import ProjectRepository
|
from packages.ports.project_repository import ProjectRepository
|
||||||
from packages.shared.storage import SharedStorageService
|
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,空为默认自然",
|
description="情绪:neutral/happy/sad/angry/surprised/fearful/disgusted,兼容旧值 natural/excited/calm/friendly,空为默认自然",
|
||||||
),
|
),
|
||||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||||
|
db: Session = Depends(get_db_session),
|
||||||
repository: SQLAlchemyVoiceCloneProfileRepository = Depends(get_voice_clone_profile_repository),
|
repository: SQLAlchemyVoiceCloneProfileRepository = Depends(get_voice_clone_profile_repository),
|
||||||
cosyvoice: CosyVoiceService = Depends(get_cosyvoice_service),
|
cosyvoice: CosyVoiceService = Depends(get_cosyvoice_service),
|
||||||
) -> VoiceClonePreviewResponse:
|
) -> VoiceClonePreviewResponse:
|
||||||
@@ -350,6 +360,31 @@ def get_voice_clone_preview(
|
|||||||
"""
|
"""
|
||||||
import time
|
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:
|
if emotion not in _ALLOWED_PREVIEW_EMOTIONS:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
@@ -393,9 +428,14 @@ def get_voice_clone_preview(
|
|||||||
speed=speed,
|
speed=speed,
|
||||||
emotion=emotion,
|
emotion=emotion,
|
||||||
)
|
)
|
||||||
except CosyVoiceError as e:
|
except (CosyVoiceError, ValueError) as e:
|
||||||
raise HTTPException(status_code=502, detail=f"TTS 合成失败: {e}") from e
|
if _points_deducted > 0 and _points_svc is not None:
|
||||||
except ValueError as e:
|
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
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
|
||||||
|
|
||||||
# 缓存(仅默认参数组合)
|
# 缓存(仅默认参数组合)
|
||||||
|
|||||||
@@ -91,13 +91,16 @@ def points_gate(
|
|||||||
is_async = asyncio.iscoroutinefunction(func)
|
is_async = asyncio.iscoroutinefunction(func)
|
||||||
|
|
||||||
if is_async:
|
if is_async:
|
||||||
|
|
||||||
async def wrapper(*args: Any, **kwargs: Any) -> Any:
|
async def wrapper(*args: Any, **kwargs: Any) -> Any:
|
||||||
if not _pg_enabled(): # noqa: F821
|
if not _pg_enabled(): # noqa: F821
|
||||||
return await func(*args, **kwargs)
|
return await func(*args, **kwargs)
|
||||||
return await _pg_execute( # noqa: F821
|
return await _pg_execute( # noqa: F821
|
||||||
func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, True
|
func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, True
|
||||||
)
|
)
|
||||||
|
|
||||||
else:
|
else:
|
||||||
|
|
||||||
def wrapper(*args: Any, **kwargs: Any) -> Any:
|
def wrapper(*args: Any, **kwargs: Any) -> Any:
|
||||||
if not _pg_enabled(): # noqa: F821
|
if not _pg_enabled(): # noqa: F821
|
||||||
return func(*args, **_pg_filter(func, kwargs)) # noqa: F821
|
return func(*args, **_pg_filter(func, kwargs)) # noqa: F821
|
||||||
@@ -126,8 +129,7 @@ def points_gate(
|
|||||||
# 为稳妥起见,直接把 wrapper code 的 co_names 映射到新名——复杂度过高,
|
# 为稳妥起见,直接把 wrapper code 的 co_names 映射到新名——复杂度过高,
|
||||||
# 这里采用「确保短名没冲突」策略:如果冲突就抛异常让开发者改名。
|
# 这里采用「确保短名没冲突」策略:如果冲突就抛异常让开发者改名。
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"points_gate: name collision in {func.__module__}.{func.__name__}: "
|
f"points_gate: name collision in {func.__module__}.{func.__name__}: " f"'{k}' already defined"
|
||||||
f"'{k}' already defined"
|
|
||||||
)
|
)
|
||||||
merged_globals[k] = v
|
merged_globals[k] = v
|
||||||
|
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -10,6 +10,16 @@ from unittest.mock import MagicMock, patch
|
|||||||
|
|
||||||
import pytest
|
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")
|
os.environ.setdefault("JWT_SECRET_KEY", "dev-secret-key-for-testing")
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -5,6 +5,15 @@ from __future__ import annotations
|
|||||||
import pytest
|
import pytest
|
||||||
from pydantic import ValidationError
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
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 test_generation_cover_router_importable():
|
def test_generation_cover_router_importable():
|
||||||
"""新路由模块可以正确导入"""
|
"""新路由模块可以正确导入"""
|
||||||
|
|||||||
@@ -0,0 +1,28 @@
|
|||||||
|
"""AI封面生成 积分扣点单元测试 (#1895 P2 step 2.7)"""
|
||||||
|
|
||||||
|
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 TestGenerationCoverPoints:
|
||||||
|
def test_ai_cover_cost(self):
|
||||||
|
from packages.domain.points_rules import calculate_points_cost
|
||||||
|
|
||||||
|
assert calculate_points_cost("ai_cover", is_member=False) == 2
|
||||||
|
assert calculate_points_cost("ai_cover", is_member=True, member_type="yearly") >= 0
|
||||||
|
|
||||||
|
def test_decorator_attached(self):
|
||||||
|
from app.api.routes.generation_cover import generate_cover
|
||||||
|
|
||||||
|
assert hasattr(generate_cover, "__wrapped__"), "missing @points_gate"
|
||||||
@@ -573,6 +573,15 @@ from app.schemas.generation_task import (
|
|||||||
PreviewGenerationTaskResponse,
|
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"):
|
def _make_user(user_id="test_user_001"):
|
||||||
"""构造 mock AuthenticatedUser"""
|
"""构造 mock AuthenticatedUser"""
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -6,6 +6,7 @@ from unittest.mock import MagicMock
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
import packages.middleware.points_gate as _pg_module
|
||||||
from packages.application.generation_tasks import (
|
from packages.application.generation_tasks import (
|
||||||
CreateGenerationTaskCommand,
|
CreateGenerationTaskCommand,
|
||||||
CreateGenerationTaskUseCase,
|
CreateGenerationTaskUseCase,
|
||||||
@@ -18,6 +19,13 @@ from packages.application.generation_tasks import (
|
|||||||
from packages.domain import GenerationTask
|
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
|
@pytest.fixture
|
||||||
def mock_repo():
|
def mock_repo():
|
||||||
return MagicMock()
|
return MagicMock()
|
||||||
|
|||||||
@@ -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"
|
||||||
@@ -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)
|
||||||
@@ -16,6 +16,8 @@ from unittest.mock import MagicMock, patch
|
|||||||
import pydantic
|
import pydantic
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
import packages.middleware.points_gate as _pg_module
|
||||||
|
|
||||||
sys.path.insert(0, "apps/api")
|
sys.path.insert(0, "apps/api")
|
||||||
|
|
||||||
|
|
||||||
@@ -28,6 +30,18 @@ def _make_auth_user(user_id: str = "u1"):
|
|||||||
return auth
|
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(
|
def _mock_youtube_dl(
|
||||||
extract_info_return=None,
|
extract_info_return=None,
|
||||||
extract_info_side_effect=None,
|
extract_info_side_effect=None,
|
||||||
@@ -82,7 +96,7 @@ class TestExtractFromDouyin:
|
|||||||
|
|
||||||
req = ExtractFromDouyinRequest(url="https://v.douyin.com/xxxxx/")
|
req = ExtractFromDouyinRequest(url="https://v.douyin.com/xxxxx/")
|
||||||
auth = _make_auth_user()
|
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.text == "这是一段测试文案内容"
|
||||||
assert result.duration_seconds == 120.5
|
assert result.duration_seconds == 120.5
|
||||||
@@ -112,7 +126,7 @@ class TestExtractFromDouyin:
|
|||||||
auth = _make_auth_user()
|
auth = _make_auth_user()
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
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
|
assert exc_info.value.status_code == 400
|
||||||
|
|
||||||
@patch("tempfile.TemporaryDirectory")
|
@patch("tempfile.TemporaryDirectory")
|
||||||
@@ -136,7 +150,7 @@ class TestExtractFromDouyin:
|
|||||||
auth = _make_auth_user()
|
auth = _make_auth_user()
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
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
|
assert exc_info.value.status_code == 502
|
||||||
|
|
||||||
@patch("app.api.routes.scripts_ai.transcribe_to_text")
|
@patch("app.api.routes.scripts_ai.transcribe_to_text")
|
||||||
@@ -169,7 +183,7 @@ class TestExtractFromDouyin:
|
|||||||
auth = _make_auth_user()
|
auth = _make_auth_user()
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
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
|
assert exc_info.value.status_code == 503
|
||||||
|
|
||||||
@patch("app.api.routes.scripts_ai.transcribe_to_text")
|
@patch("app.api.routes.scripts_ai.transcribe_to_text")
|
||||||
@@ -202,7 +216,7 @@ class TestExtractFromDouyin:
|
|||||||
auth = _make_auth_user()
|
auth = _make_auth_user()
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
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
|
assert exc_info.value.status_code == 502
|
||||||
|
|
||||||
|
|
||||||
@@ -225,7 +239,7 @@ class TestAiRewrite:
|
|||||||
|
|
||||||
req = AiRewriteRequest(content="原始文案内容", style="口语化")
|
req = AiRewriteRequest(content="原始文案内容", style="口语化")
|
||||||
auth = _make_auth_user()
|
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.original == "原始文案内容"
|
||||||
assert result.rewritten == "改写后的文案内容,口语化风格"
|
assert result.rewritten == "改写后的文案内容,口语化风格"
|
||||||
@@ -242,7 +256,7 @@ class TestAiRewrite:
|
|||||||
auth = _make_auth_user()
|
auth = _make_auth_user()
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
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
|
assert exc_info.value.status_code == 400
|
||||||
|
|
||||||
@patch("app.api.routes.scripts_ai.get_doubao_client")
|
@patch("app.api.routes.scripts_ai.get_doubao_client")
|
||||||
@@ -261,7 +275,7 @@ class TestAiRewrite:
|
|||||||
auth = _make_auth_user()
|
auth = _make_auth_user()
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
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
|
assert exc_info.value.status_code == 502
|
||||||
|
|
||||||
@patch("app.api.routes.scripts_ai.get_doubao_client")
|
@patch("app.api.routes.scripts_ai.get_doubao_client")
|
||||||
@@ -279,7 +293,7 @@ class TestAiRewrite:
|
|||||||
auth = _make_auth_user()
|
auth = _make_auth_user()
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
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
|
assert exc_info.value.status_code == 502
|
||||||
|
|
||||||
|
|
||||||
@@ -301,7 +315,7 @@ class TestAiGenerateTitles:
|
|||||||
|
|
||||||
req = AiGenerateTitlesRequest(content="这是一段关于美食的文案", count=3)
|
req = AiGenerateTitlesRequest(content="这是一段关于美食的文案", count=3)
|
||||||
auth = _make_auth_user()
|
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 len(result.titles) == 3
|
||||||
assert all(isinstance(t, str) for t in result.titles)
|
assert all(isinstance(t, str) for t in result.titles)
|
||||||
@@ -339,12 +353,12 @@ class TestAiGenerateTitles:
|
|||||||
|
|
||||||
# count=5
|
# count=5
|
||||||
req = AiGenerateTitlesRequest(content="测试内容", 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
|
assert len(result.titles) <= 5
|
||||||
|
|
||||||
# count=1
|
# count=1
|
||||||
req = AiGenerateTitlesRequest(content="测试内容", 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
|
assert len(result.titles) >= 1
|
||||||
|
|
||||||
def test_generate_titles_empty_content(self):
|
def test_generate_titles_empty_content(self):
|
||||||
@@ -357,7 +371,7 @@ class TestAiGenerateTitles:
|
|||||||
auth = _make_auth_user()
|
auth = _make_auth_user()
|
||||||
|
|
||||||
with pytest.raises(HTTPException) as exc_info:
|
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
|
assert exc_info.value.status_code == 400
|
||||||
|
|
||||||
@patch("app.services.ai_service.get_doubao_client")
|
@patch("app.services.ai_service.get_doubao_client")
|
||||||
@@ -372,7 +386,7 @@ class TestAiGenerateTitles:
|
|||||||
|
|
||||||
req = AiGenerateTitlesRequest(content="测试文案内容")
|
req = AiGenerateTitlesRequest(content="测试文案内容")
|
||||||
auth = _make_auth_user()
|
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 len(result.titles) == 3
|
||||||
|
|
||||||
|
|||||||
@@ -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)
|
||||||
@@ -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")
|
||||||
Reference in New Issue
Block a user