feat(#1895): 会员积分系统后端(P1-P4,默认关闭) #1923

Closed
xiaoxia wants to merge 7 commits from feat/1895-membership-points-p1-p2 into develop
28 changed files with 2503 additions and 1141 deletions
+78 -66
View File
@@ -1,12 +1,11 @@
"""add membership & points system
"""membership + points tables
Revision ID: 076_membership_points
Revises: 075_add_sentence_timings
Create Date: 2026-09-15
Create Date: 2026-09-14
"""
import sqlalchemy as sa
from sqlalchemy import text
from alembic import op
@@ -17,100 +16,101 @@ depends_on = None
def upgrade() -> None:
# 1. users 表新增字段
# 1. users 表加字段
with op.batch_alter_table("users") as batch:
batch.add_column(
sa.Column("is_member", sa.Boolean(), nullable=False, server_default=sa.text("false")),
sa.Column("is_member", sa.Boolean(), nullable=False, server_default=sa.false()),
)
batch.add_column(sa.Column("member_type", sa.String(length=20), nullable=True))
batch.add_column(sa.Column("member_expires_at", sa.DateTime(), nullable=True))
batch.add_column(
sa.Column("member_type", sa.String(20), nullable=True),
)
batch.add_column(
sa.Column("member_expires_at", sa.DateTime(), nullable=True),
)
batch.add_column(
sa.Column("points_balance", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column(
"points_balance",
sa.Integer(),
nullable=False,
server_default=sa.text("0"),
),
)
# 2. points_accounts 积分账户表
# 2. points_accounts 表
op.create_table(
"points_accounts",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, unique=True, index=True),
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("user_id", sa.String(length=36), nullable=False),
sa.Column("balance", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("total_earned", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("total_spent", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
sa.Column(
"updated_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
sa.Column("total_purchased", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("total_gifted", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
sa.Column("updated_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint("user_id", name="uq_points_accounts_user_id"),
)
op.create_index("idx_points_accounts_user", "points_accounts", ["user_id"])
# 3. points_transactions 积分流水表
# 3. points_transactions 表
op.create_table(
"points_transactions",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("account_id", sa.String(36), nullable=False, index=True),
sa.Column("type", sa.String(20), nullable=False, index=True),
sa.Column("source", sa.String(50), nullable=False, index=True),
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("user_id", sa.String(length=36), nullable=False),
sa.Column("account_id", sa.String(length=36), nullable=False),
sa.Column("type", sa.String(length=20), nullable=False),
sa.Column("source", sa.String(length=50), nullable=False),
sa.Column("amount", sa.Integer(), nullable=False),
sa.Column("balance_after", sa.Integer(), nullable=False),
sa.Column("description", sa.String(255), nullable=False, server_default=""),
sa.Column("ref_id", sa.String(100), nullable=False, server_default=""),
sa.Column(
"created_at",
sa.DateTime(),
"description",
sa.String(length=255),
nullable=False,
server_default=sa.text("NOW()"),
server_default="",
),
sa.Column("ref_id", sa.String(length=100), nullable=False, server_default=""),
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
sa.ForeignKeyConstraint(["account_id"], ["points_accounts.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("idx_points_tx_user", "points_transactions", ["user_id"])
op.create_index("idx_points_tx_type", "points_transactions", ["type"])
op.create_index("idx_points_tx_source", "points_transactions", ["source"])
op.create_index("idx_points_tx_created", "points_transactions", ["created_at"])
# 4. points_orders 积分/会员订单表
# 4. points_orders 表
op.create_table(
"points_orders",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("order_type", sa.String(20), nullable=False),
sa.Column("product_code", sa.String(50), nullable=False),
sa.Column("amount_cents", sa.Integer(), nullable=False),
sa.Column("original_amount_cents", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("user_id", sa.String(length=36), nullable=False),
sa.Column("package_name", sa.String(length=50), nullable=False),
sa.Column("points_amount", sa.Integer(), nullable=False),
sa.Column("price_cents", sa.Integer(), nullable=False),
sa.Column("currency", sa.String(length=10), nullable=False, server_default="CNY"),
sa.Column("discount", sa.Float(), nullable=False, server_default=sa.text("1.0")),
sa.Column("points_amount", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column("status", sa.String(20), nullable=False, server_default="pending", index=True),
sa.Column("payment_method", sa.String(50), nullable=True),
sa.Column("payment_id", sa.String(100), nullable=True),
sa.Column("original_price_cents", sa.Integer(), nullable=False),
sa.Column("status", sa.String(length=20), nullable=False, server_default="pending"),
sa.Column("payment_method", sa.String(length=50), nullable=True),
sa.Column("payment_id", sa.String(length=100), nullable=True),
sa.Column("paid_at", sa.DateTime(), nullable=True),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
sa.Column("expire_at", sa.DateTime(), nullable=True),
sa.Column("created_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
)
op.create_index("idx_points_orders_user", "points_orders", ["user_id"])
op.create_index("idx_points_orders_status", "points_orders", ["status"])
# 5. daily_usage_records 每日使用记录表
# 5. daily_usage_records 表
op.create_table(
"daily_usage_records",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False, index=True),
sa.Column("usage_date", sa.DateTime(), nullable=False),
sa.Column("usage_type", sa.String(50), nullable=False, server_default="free_clip"),
sa.Column("id", sa.String(length=36), nullable=False),
sa.Column("user_id", sa.String(length=36), nullable=False),
sa.Column("usage_date", sa.Date(), nullable=False),
sa.Column("usage_type", sa.String(length=50), nullable=False),
sa.Column("count", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column(
"updated_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
sa.Column("updated_at", sa.DateTime(), nullable=False, server_default=sa.func.now()),
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
sa.PrimaryKeyConstraint("id"),
sa.UniqueConstraint(
"user_id",
"usage_date",
@@ -118,12 +118,24 @@ def upgrade() -> None:
name="uq_daily_usage_user_date_type",
),
)
op.create_index("idx_daily_usage_user_date", "daily_usage_records", ["user_id", "usage_date"])
def downgrade() -> None:
op.drop_index("idx_daily_usage_user_date", table_name="daily_usage_records")
op.drop_table("daily_usage_records")
op.drop_index("idx_points_orders_status", table_name="points_orders")
op.drop_index("idx_points_orders_user", table_name="points_orders")
op.drop_table("points_orders")
op.drop_index("idx_points_tx_created", table_name="points_transactions")
op.drop_index("idx_points_tx_source", table_name="points_transactions")
op.drop_index("idx_points_tx_type", table_name="points_transactions")
op.drop_index("idx_points_tx_user", table_name="points_transactions")
op.drop_table("points_transactions")
op.drop_index("idx_points_accounts_user", table_name="points_accounts")
op.drop_table("points_accounts")
with op.batch_alter_table("users") as batch:
+12 -16
View File
@@ -6,7 +6,6 @@ from app.api.routes.assets import router as assets_router
from app.api.routes.auth import router as auth_router
from app.api.routes.chunked_upload import router as chunked_upload_router
from app.api.routes.classification_jobs import router as classification_jobs_router
from app.api.routes.clips_standalone import router as clips_standalone_router
from app.api.routes.cover_templates import router as cover_templates_router
from app.api.routes.duplication import router as duplication_router
from app.api.routes.feature_flags import router as feature_flags_router
@@ -18,7 +17,7 @@ from app.api.routes.health import router as health_check_router
from app.api.routes.ingest_jobs import router as ingest_jobs_router
from app.api.routes.internal_render import router as internal_render_router
from app.api.routes.lipsync import router as lipsync_router
from app.api.routes.points import points_router, usage_router
from app.api.routes.points import router as points_router
from app.api.routes.projects import router as projects_router
from app.api.routes.scripts import router as scripts_router
from app.api.routes.share import router as share_router
@@ -30,6 +29,7 @@ from app.api.routes.templates_editor import router as templates_editor_router
from app.api.routes.titles import router as titles_router
from app.api.routes.tts import router as tts_router
from app.api.routes.upload import router as upload_router
from app.api.routes.usage import router as usage_router
from app.api.routes.videos import router as videos_router
from app.api.routes.voice_clones import router as voice_clones_router
from app.api.routes.voices import router as voices_router
@@ -153,15 +153,21 @@ api_router.include_router(
prefix="/subscription",
tags=["Subscription"],
)
api_router.include_router(
points_router,
prefix="/points",
tags=["Points"],
)
api_router.include_router(
usage_router,
prefix="/usage",
tags=["Usage"],
)
api_router.include_router(
templates_router,
prefix="/templates",
tags=["Template"],
)
api_router.include_router(
clips_standalone_router,
tags=["Clips"],
)
api_router.include_router(
templates_editor_router,
prefix="/templates/{template_id}/editor",
@@ -195,13 +201,3 @@ api_router.include_router(
prefix="/ai-avatar/render",
tags=["AI Avatar Render"],
)
api_router.include_router(
points_router,
prefix="/points",
tags=["Points"],
)
api_router.include_router(
usage_router,
prefix="/usage",
tags=["Usage"],
)
+22 -7
View File
@@ -7,10 +7,15 @@ from __future__ import annotations
from typing import List, Literal
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_points_service
from app.services.ai_service import TITLE_STYLES, generate_smart_titles, semantic_match_assets
from fastapi import APIRouter
from fastapi import APIRouter, Depends
from pydantic import BaseModel, Field
from packages.application.points_service import PointsService
from packages.middleware.points_gate import points_deduction
router = APIRouter()
@@ -85,17 +90,27 @@ class SemanticMatchResponse(BaseModel):
@router.post("/titles/generate", response_model=GenerateTitlesResponse)
def generate_titles(request: GenerateTitlesRequest):
def generate_titles(
request: GenerateTitlesRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
points_svc: PointsService = Depends(get_points_service),
):
"""生成智能标题.
根据视频描述生成指定风格的标题,支持爆款、情感、信息三种风格。
未配置豆包 API Key 时自动降级为本地规则生成。
"""
result = generate_smart_titles(
description=request.description,
style=request.style,
count=request.count,
)
with points_deduction(
points_svc,
authenticated_user.user,
"ai_title",
description="AI 标题生成",
):
result = generate_smart_titles(
description=request.description,
style=request.style,
count=request.count,
)
return GenerateTitlesResponse(**result)
+37 -1
View File
@@ -14,7 +14,8 @@ import logging
from datetime import datetime, timezone
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.config import settings
from app.dependencies import get_db_session, get_points_service
from app.schemas.ai_avatar_render import (
AiAvatarRenderJobResponse,
CreateAiAvatarRenderRequest,
@@ -29,6 +30,9 @@ from app.services.ai_avatar_render_service import (
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from packages.application.points_service import PointsService
from packages.middleware.points_gate import _insufficient_points, _is_active_member
logger = logging.getLogger(__name__)
router = APIRouter()
@@ -46,11 +50,34 @@ def create_render_job(
body: CreateAiAvatarRenderRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: AiAvatarRenderService = Depends(_get_service),
points_svc: PointsService = Depends(get_points_service),
):
"""提交 AI 数字人渲染任务.
将对口型视频 + B-roll 素材 + 标题叠加 + 封面提取合成最终输出视频。
"""
_pts_tx_id: str | None = None
_user = current_user.user
if settings.POINTS_ENABLED:
_is_m = _is_active_member(_user)
_dr = points_svc.check_and_deduct(
user_id=_user.id,
scene_key="ai_digital_human",
duration_minutes=1,
description="AI 数字人渲染",
is_member=_is_m,
)
if not _dr.success:
raise _insufficient_points(_dr.amount, _dr.balance, "AI 数字人渲染")
_pts_tx_id = _dr.transaction_id
def _refund(reason: str) -> None:
if _pts_tx_id:
try:
points_svc.refund(_user.id, _pts_tx_id, reason=reason)
except Exception:
logger.exception("数字人渲染退款失败")
try:
job = svc.create_render_job(
user_id=current_user.user.id,
@@ -62,6 +89,7 @@ def create_render_job(
project_id=body.project_id,
)
except AiAvatarRenderError as exc:
_refund(f"数字人渲染业务错误: {exc.code}")
status_map = {
"LipsyncJobNotFound": 404,
"LipsyncJobNotCompleted": 400,
@@ -72,6 +100,12 @@ def create_render_job(
status_code=status_map.get(exc.code, 400),
detail={"code": exc.code, "message": str(exc)},
) from exc
except HTTPException:
_refund("数字人渲染HTTP异常")
raise
except Exception:
_refund("数字人渲染异常")
raise
# 异步触发渲染
try:
@@ -80,6 +114,7 @@ def create_render_job(
execute_ai_avatar_render.delay(job.id)
except Exception as exc:
logger.exception("Celery 任务投递失败(创建): job_id=%s err=%s", job.id, exc)
_refund("数字人渲染Celery投递失败")
job.status = "failed"
job.error_message = f"任务提交失败:{exc}"
job.updated_at = datetime.now(timezone.utc)
@@ -259,6 +294,7 @@ def generate_render_smart_cover(
)
return SmartCoverResponse(cover_url=cover_url, status="completed")
# ── POST /{job_id}/finalize — 封面选定后正式入库成片库 ────────────────────
+16 -1
View File
@@ -13,7 +13,7 @@ from typing import Optional
import jwt
from app.auth import AuthenticatedUser, blacklist_token, get_current_user
from app.config import settings
from app.dependencies import get_auth_email_service, get_auth_session_store, get_user_repository
from app.dependencies import get_auth_email_service, get_auth_session_store, get_points_service, get_user_repository
from fastapi import APIRouter, Depends, Header, HTTPException, Request, status
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from pydantic import BaseModel, EmailStr, field_validator
@@ -32,6 +32,8 @@ from packages.application.auth.password_reset_use_case import (
)
from packages.application.auth.register_user_use_case import RegisterUserRequest as RegisterUseCaseRequest
from packages.application.auth.register_user_use_case import RegisterUserUseCase, VerifyEmailRequest, VerifyEmailUseCase
from packages.application.points_service import PointsService
from packages.domain.points import TX_SOURCE_TASK_REWARD
from packages.ports.user_repository import UserRepository
logger = logging.getLogger(__name__)
@@ -126,6 +128,7 @@ async def register(
request: RegisterRequest,
user_repository: UserRepository = Depends(get_user_repository),
email_service=Depends(get_auth_email_service),
points_svc: PointsService = Depends(get_points_service),
) -> RegisterResponse:
use_case = RegisterUserUseCase(
user_repository=user_repository,
@@ -143,6 +146,18 @@ async def register(
if error or response is None:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=_translate_auth_error(error))
# 新用户注册送 50 积分(#1895 P4),失败不影响注册
if settings.POINTS_ENABLED:
try:
points_svc.earn_points(
user_id=response.user_id,
amount=50,
source=TX_SOURCE_TASK_REWARD,
description="新用户注册赠送",
)
except Exception:
logging.getLogger(__name__).exception("新用户注册送积分失败 user=%s", response.user_id)
return RegisterResponse(
user_id=response.user_id,
email=response.email,
+41 -5
View File
@@ -15,7 +15,8 @@ from typing import Any, List, Optional
from urllib.parse import urlparse
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session, get_generated_video_repository
from app.config import settings
from app.dependencies import get_db_session, get_generated_video_repository, get_points_service
from app.services.edit_plan_service import EditPlanService
from app.services.edit_template_service import EditTemplateService
from fastapi import APIRouter, Depends, HTTPException, Query
@@ -26,7 +27,9 @@ from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.application import ListGeneratedVideosByTaskUseCase
from packages.application.points_service import PointsService
from packages.domain.config_schemas import normalize_plan_config
from packages.middleware.points_gate import _insufficient_points, _is_active_member
from packages.shared.storage import get_shared_storage_service
from .templates_editor.dependencies import get_draft_plan_id, get_editor_services
@@ -75,10 +78,7 @@ class GenerateCoverResponse(BaseModel):
# ── Route ────────────────────────────────────────────────────────────────
def _select_best_frame_from_snapshots(
snapshots: list[dict], plan_id: str
) -> str:
def _select_best_frame_from_snapshots(snapshots: list[dict], plan_id: str) -> str:
"""从 MediaKit 抽帧结果中,通过质量评分选出最佳帧。
降级策略:cv2 不可用或评分失败时,返回第一帧。
@@ -338,6 +338,7 @@ def generate_cover(
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
db: Session = Depends(get_db_session),
current_user: AuthenticatedUser = Depends(get_current_user),
points_svc: PointsService = Depends(get_points_service),
) -> GenerateCoverResponse:
"""AI 生成封面 — 优先从最终成片视频中抽帧,回退到预览片段.
@@ -856,6 +857,22 @@ def generate_cover(
from packages.shared.ai_service import run_generate_cover
# 积分扣费(#1895 P4):AI 封面 1 积分/张;upload 类型已提前 return 不走这里
_pts_tx_id: str | None = None
_user = current_user.user
if settings.POINTS_ENABLED:
_is_m = _is_active_member(_user)
_dr = points_svc.check_and_deduct(
user_id=_user.id,
scene_key="ai_cover",
duration_minutes=1,
description="AI 封面生成",
is_member=_is_m,
)
if not _dr.success:
raise _insufficient_points(_dr.amount, _dr.balance, "AI 封面生成")
_pts_tx_id = _dr.transaction_id
try:
logger.info("[封面生成] 开始调用 AI 封面生成服务: plan_id=%s", plan_id)
cover_data = run_generate_cover(
@@ -866,7 +883,26 @@ def generate_cover(
primary_video_url=primary_video_url,
)
except RuntimeError as e:
if _pts_tx_id:
try:
points_svc.refund(_user.id, _pts_tx_id, reason="AI封面 RuntimeError")
except Exception:
logger.exception("AI封面退款失败")
raise HTTPException(status_code=500, detail=str(e)) from e
except HTTPException:
if _pts_tx_id:
try:
points_svc.refund(_user.id, _pts_tx_id, reason="AI封面 HTTP异常")
except Exception:
logger.exception("AI封面退款失败")
raise
except Exception:
if _pts_tx_id:
try:
points_svc.refund(_user.id, _pts_tx_id, reason="AI封面异常")
except Exception:
logger.exception("AI封面退款失败")
raise
current_config = dict(plan.config) if plan.config else {}
current_config["cover"] = cover_data
@@ -8,6 +8,7 @@ from __future__ import annotations
import logging
from app.auth import AuthenticatedUser, get_current_user
from app.config import settings
from app.core.storage import get_storage_service
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
@@ -22,6 +23,7 @@ from app.dependencies import (
get_db_session,
get_generated_video_repository,
get_generation_task_repository,
get_points_service,
)
from app.schemas.generation_task import (
BatchPreviewGenerationTaskResponse,
@@ -43,6 +45,8 @@ from packages.application import (
GetGenerationTaskUseCase,
ListGeneratedVideosByTaskUseCase,
)
from packages.application.points_service import PointsService
from packages.middleware.points_gate import _insufficient_points, _is_active_member
logger = logging.getLogger(__name__)
@@ -277,6 +281,7 @@ def create_preview_generation_task(
generation_task_repository=Depends(get_generation_task_repository),
db: Session = Depends(get_db_session),
asset_repo=Depends(get_asset_repository),
points_svc: PointsService = Depends(get_points_service),
) -> BatchPreviewGenerationTaskResponse:
"""创建预览生成任务(支持批量)。
@@ -292,6 +297,48 @@ def create_preview_generation_task(
"""
user_id = authenticated_user.user.id
count = max(1, request.preview_count)
_pts_enabled: bool = bool(settings.POINTS_ENABLED)
_pts_is_member: bool = _is_active_member(authenticated_user.user) if _pts_enabled else False
_pts_tx_ids: list[str] = []
def _pts_refund_all(reason: str) -> None:
for _tx in list(_pts_tx_ids):
try:
points_svc.refund(user_id, _tx, reason=reason)
except Exception:
logger.exception("预览生成退款失败 user=%s tx=%s", user_id, _tx)
_pts_tx_ids.clear()
def _pts_refund_last(reason: str) -> None:
if _pts_tx_ids:
_tx = _pts_tx_ids.pop()
try:
points_svc.refund(user_id, _tx, reason=reason)
except Exception:
logger.exception("预览单条退款失败 user=%s tx=%s", user_id, _tx)
def _pts_deduct_one(idx: int) -> None:
if not _pts_enabled:
return
if not _pts_is_member:
try:
if points_svc.check_and_incr_daily_free_clips(user_id):
return
except Exception:
logger.warning("[预览生成] daily_free_clips 异常,降级走扣费 user=%s", user_id, exc_info=True)
_dr = points_svc.check_and_deduct(
user_id=user_id,
scene_key="ai_video",
duration_minutes=1,
description=f"智能混剪预览(第{idx+1}条)",
is_member=_pts_is_member,
)
if not _dr.success:
_pts_refund_all("预览扣费失败回退")
raise _insufficient_points(_dr.amount, _dr.balance, "智能混剪预览")
if _dr.transaction_id:
_pts_tx_ids.append(_dr.transaction_id)
logger.info(
"[预览生成] 接收请求: user_id=%s, template_id=%s, asset_count=%d, preview_count=%d",
user_id,
@@ -407,6 +454,8 @@ def create_preview_generation_task(
)
task.extra_meta["variant_index"] = variant_index
_pts_deduct_one(variant_index)
# 解析源编辑计划(前端传入或按模板兜底查找)
source_plan_id = _resolve_preview_edit_plan_id(request=request, task=task, db=db, user_id=user_id)
task.source_edit_plan_id = source_plan_id
@@ -414,9 +463,14 @@ def create_preview_generation_task(
created_tasks.append(task)
except ValueError as e:
logger.warning("[预览生成] 创建失败: %s", e)
_pts_refund_all("预览创建失败(ValueError)")
raise HTTPException(status_code=400, detail=str(e)) from e
except HTTPException:
_pts_refund_last("预览创建参数异常")
raise
except Exception as e:
logger.error("[预览生成] 创建失败: %s", e, exc_info=True)
_pts_refund_all("预览创建失败(异常)")
raise HTTPException(status_code=500, detail="创建预览生成任务失败,请稍后再试") from e
# ── 独立变体 plan(#1743)──
@@ -539,9 +593,11 @@ def create_preview_generation_task(
) from last_err
variant_plan_ids.append(variant_plan.id)
except HTTPException:
_pts_refund_all("预览变体参数异常")
raise
except Exception as e:
logger.error("[预览生成] 变体 plan 生成异常: %s", e, exc_info=True)
_pts_refund_all("预览变体 plan 异常")
for t in created_tasks:
_mark_task_failed(generation_task_repository, t, "预览变体计划创建失败")
raise HTTPException(
@@ -575,6 +631,7 @@ def create_preview_generation_task(
# ── 入队 ──
responses: list[PreviewGenerationTaskResponse] = []
rate_limit_exc: Exception | None = None # 记录首个限流异常,全部失败时返回结构化提示
_enqueued_count = 0
for variant_index, task in enumerate(created_tasks):
try:
enqueued = safe_enqueue_generation_task(
@@ -587,17 +644,25 @@ def create_preview_generation_task(
if not enqueued:
logger.warning("[预览生成] 任务入队失败: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队失败")
_pts_refund_last("预览任务入队失败")
else:
_enqueued_count += 1
except UserPendingLimitExceeded as e:
_mark_task_failed(generation_task_repository, task, "待处理任务超限")
_pts_refund_last("预览用户限流")
rate_limit_exc = rate_limit_exc or e
except GlobalQueueFull as e:
_mark_task_failed(generation_task_repository, task, "系统队列已满")
_pts_refund_last("预览全局限流")
rate_limit_exc = rate_limit_exc or e
except Exception:
logger.exception("[预览生成] 入队异常: task_id=%s", task.id)
_mark_task_failed(generation_task_repository, task, "任务入队异常")
_pts_refund_last("预览入队异常")
# enqueue 会原地更新 task 状态/进度,直接用 task 构造响应
responses.append(_to_preview_response(task))
if _enqueued_count == 0 and _pts_tx_ids:
_pts_refund_all("预览无任务入队成功")
# 队列满/限流时若全部失败,返回结构化错误码(前端区分"排队"与"创建失败")
if all(r.status == "failed" for r in responses) and rate_limit_exc is not None:
+66 -5
View File
@@ -4,6 +4,7 @@ from typing import Any
from app.api.routes._helpers import check_project_access
from app.auth import AuthenticatedUser, get_current_user
from app.config import settings
from app.core.storage import OSSStorageService, get_storage_service
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
@@ -19,6 +20,7 @@ from app.dependencies import (
get_db_session,
get_generated_video_repository,
get_generation_task_repository,
get_points_service,
get_project_repository,
)
from app.schemas.generated_video import (
@@ -41,7 +43,9 @@ from packages.application import (
GetGenerationTaskUseCase,
ListGeneratedVideosByTaskUseCase,
)
from packages.application.points_service import PointsService
from packages.domain.smart_match import smart_select_assets
from packages.middleware.points_gate import _insufficient_points, _is_active_member
logger = logging.getLogger(__name__)
@@ -219,15 +223,59 @@ def create_generation_task(
asset_library_repository: Any = Depends(get_asset_library_repository),
asset_repository: Any = Depends(get_asset_repository),
db: Session = Depends(get_db_session),
points_svc: PointsService = Depends(get_points_service),
) -> BatchGenerationTaskResponse:
user_id = authenticated_user.user.id
logger.info(
"[生成任务] 接收请求: user_id=%s, template_id=%s, asset_count=%d, mode=%s, count=%d",
authenticated_user.user.id,
user_id,
request.template_id,
len(request.asset_ids),
request.asset_select_mode,
request.count,
)
# 积分开关 & 会员状态(#1895 P4)
_pts_enabled: bool = bool(settings.POINTS_ENABLED)
_pts_is_member: bool = _is_active_member(authenticated_user.user) if _pts_enabled else False
_pts_tx_ids: list[str] = []
def _pts_refund_all(reason: str) -> None:
for _tx in list(_pts_tx_ids):
try:
points_svc.refund(user_id, _tx, reason=reason)
except Exception:
logger.exception("批量生成退款失败 user=%s tx=%s", user_id, _tx)
_pts_tx_ids.clear()
def _pts_refund_last(reason: str) -> None:
if _pts_tx_ids:
_tx = _pts_tx_ids.pop()
try:
points_svc.refund(user_id, _tx, reason=reason)
except Exception:
logger.exception("单条任务退款失败 user=%s tx=%s", user_id, _tx)
def _pts_deduct_one(idx: int) -> None:
if not _pts_enabled:
return
if not _pts_is_member:
try:
if points_svc.check_and_incr_daily_free_clips(user_id):
return
except Exception:
logger.warning("[生成任务] daily_free_clips 异常,降级走扣费 user=%s", user_id, exc_info=True)
_dr = points_svc.check_and_deduct(
user_id=user_id,
scene_key="ai_video",
duration_minutes=1,
description=f"智能混剪(第{idx+1}条)",
is_member=_pts_is_member,
)
if not _dr.success:
_pts_refund_all("智能混剪批量扣费失败回退")
raise _insufficient_points(_dr.amount, _dr.balance, "智能混剪")
if _dr.transaction_id:
_pts_tx_ids.append(_dr.transaction_id)
try:
project_id, asset_library_id = _resolve_project_and_library(
@@ -361,7 +409,6 @@ def create_generation_task(
count = request.count
created_tasks: list = []
failed_tasks = []
user_id = authenticated_user.user.id
# 同批次任务共享 batch_id,用于视频查重时批次内比对
batch_id = uuid.uuid4().hex if count > 1 else ""
@@ -614,6 +661,8 @@ def create_generation_task(
title_config=variant_title_config,
)
)
# 积分扣费(#1895 P4)
_pts_deduct_one(task_index)
# 变体序号写入 extra_meta(响应/排查时可辨识)
task.extra_meta["variant_index"] = task_index
try:
@@ -672,20 +721,23 @@ def create_generation_task(
db=db,
)
if safe_enqueue_generation_task(
_enq_ok = safe_enqueue_generation_task(
task,
generation_task_repository,
user_id=user_id,
log_prefix="[生成任务]",
log_task_status=True,
):
)
if _enq_ok:
created_tasks.append(task)
else:
failed_tasks.append(task)
_pts_refund_last("智能混剪入队失败")
except UserPendingLimitExceeded as _e:
# 兜底:如果预检查后又并发提交了,在这里也拦住
failed_tasks.append(task)
_pts_refund_last("智能混剪用户限流")
if not created_tasks:
_pts_refund_all("智能混剪用户限流(全部)")
raise HTTPException(
status_code=429,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="user"),
@@ -693,18 +745,27 @@ def create_generation_task(
break
except GlobalQueueFull as _e:
failed_tasks.append(task)
_pts_refund_last("智能混剪全局限流")
if not created_tasks:
_pts_refund_all("智能混剪全局限流(全部)")
raise HTTPException(
status_code=503,
detail=build_rate_limit_detail(_e, generation_task_repository, scope="global"),
) from _e
break
except HTTPException:
_pts_refund_last("智能混剪参数异常")
raise
except HTTPException:
raise
except Exception as e:
_pts_refund_all("智能混剪异常")
logger.error("[生成任务] 创建失败: %s", e, exc_info=True)
raise HTTPException(status_code=500, detail="创建生成任务失败,请稍后重试或查看任务日志") from e
if not created_tasks and _pts_tx_ids:
_pts_refund_all("智能混剪无成功任务")
items = [_to_generation_task_response(t) for t in created_tasks + failed_tasks]
return BatchGenerationTaskResponse(items=items, total=len(items))
+43
View File
@@ -14,8 +14,10 @@ from __future__ import annotations
import logging
from app.auth import AuthenticatedUser, get_current_user
from app.config import settings
from app.dependencies import (
get_db_session,
get_points_service,
get_voice_clone_profile_repository,
)
from app.schemas.lipsync import (
@@ -29,6 +31,9 @@ from app.services.mediakit_client import MediaKitError
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from packages.application.points_service import PointsService
from packages.middleware.points_gate import _insufficient_points, _is_active_member
logger = logging.getLogger(__name__)
router = APIRouter()
@@ -53,6 +58,7 @@ def create_lipsync_job(
body: CreateLipsyncJobRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: LipsyncService = Depends(_get_service),
points_svc: PointsService = Depends(get_points_service),
):
"""提交对口型任务.
@@ -63,6 +69,21 @@ def create_lipsync_job(
- 预合成音频(#1845 新主路径):传 {video_url, audio_url, audio_duration, sentence_timings},
后端同步ffprobe+写入timings+直接提交MediaKit(~2-3s)。
"""
_pts_tx_id: str | None = None
_user = current_user.user
if settings.POINTS_ENABLED:
_is_m = _is_active_member(_user)
_dr = points_svc.check_and_deduct(
user_id=_user.id,
scene_key="ai_digital_human",
duration_minutes=1,
description="AI 对口型",
is_member=_is_m,
)
if not _dr.success:
raise _insufficient_points(_dr.amount, _dr.balance, "AI 对口型")
_pts_tx_id = _dr.transaction_id
try:
job = svc.create_job(
user_id=current_user.user.id,
@@ -78,8 +99,18 @@ def create_lipsync_job(
project_id=body.project_id,
)
except ValueError as exc:
if _pts_tx_id:
try:
points_svc.refund(_user.id, _pts_tx_id, reason="对口型参数错误")
except Exception:
logger.exception("对口型退款失败")
raise HTTPException(status_code=400, detail=str(exc)) from exc
except MediaKitError as exc:
if _pts_tx_id:
try:
points_svc.refund(_user.id, _pts_tx_id, reason=f"对口型MediaKit错误: {exc.code}")
except Exception:
logger.exception("对口型退款失败")
status_code = 502
if exc.code in ("VoiceForbidden",):
status_code = 403
@@ -93,8 +124,20 @@ def create_lipsync_job(
"request_id": getattr(exc, "request_id", ""),
},
) from exc
except HTTPException:
if _pts_tx_id:
try:
points_svc.refund(_user.id, _pts_tx_id, reason="对口型HTTP异常")
except Exception:
logger.exception("对口型退款失败")
raise
except Exception as exc:
logger.error("创建对口型任务异常: %s", exc, exc_info=True)
if _pts_tx_id:
try:
points_svc.refund(_user.id, _pts_tx_id, reason="对口型异常")
except Exception:
logger.exception("对口型退款失败")
raise HTTPException(
status_code=400,
detail=f"创建对口型任务失败: {exc}",
+244 -270
View File
@@ -1,321 +1,295 @@
"""积分 & 会员 API 路由 (#1895)
导出两个 router:
- points_router: 积分相关路由,前缀 /points
- usage_router: 每日额度路由,前缀 /usage
"""
"""积分 & 会员充值 API 路由(#1895 P3)。"""
from __future__ import annotations
import logging
from datetime import datetime
from datetime import datetime, timedelta, timezone
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.dependencies import get_points_service
from app.schemas.points import (
DailyUsageResponse,
MembershipStatusResponse,
PointRuleItem,
PointsBalanceResponse,
PointsCheckRequest,
PointsCheckResponse,
PointsDeductRequest,
PointsOrderResponse,
PointsDeductResponse,
PointsPackageItem,
PointsPackagesResponse,
PointsRechargeRequest,
PointsRechargeResponse,
PointsRefundRequest,
PointsRefundResponse,
PointsRuleItem,
PointsRulesResponse,
PointsTransactionsResponse,
SimpleMessageResponse,
PointsTransactionItem,
PointsTransactionListResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from fastapi import APIRouter, Depends, HTTPException, Query, status
from packages.domain.points_rules import (
from packages.application.points_service import (
POINTS_UNIT_PRICE_YUAN,
PointsService,
)
from packages.domain.points import (
FREE_DAILY_CLIPS,
FREE_USER_MULTIPLIER,
MEMBER_DISCOUNT,
POINTS_PACKAGES,
POINTS_SCENES,
calculate_points_cost,
MEMBER_PACKAGE_DISCOUNT,
POINTS_RULES,
)
from packages.domain.points_service import PointsService
logger = logging.getLogger(__name__)
# ── 两个 router ──
points_router = APIRouter()
usage_router = APIRouter()
router = APIRouter()
CST = timezone(timedelta(hours=8))
def _get_service() -> PointsService:
return PointsService()
def _member_discount(member_type: str | None) -> float:
if not member_type:
return 1.0
return MEMBER_PACKAGE_DISCOUNT.get(member_type, 1.0)
def _is_member(user: AuthenticatedUser) -> bool:
"""判断用户是否为付费会员。"""
return getattr(user.user, "is_member", False)
def _member_limit(user: AuthenticatedUser) -> int:
"""会员每日免费混剪条数:付费会员不限 (-1),免费用户 FREE_DAILY_CLIPS。"""
return -1 if bool(getattr(user.user, "is_member", False)) else FREE_DAILY_CLIPS
def _member_type(user: AuthenticatedUser) -> str | None:
return getattr(user.user, "member_type", None)
# ── 余额 ──────────────────────────────────────────────────────────────────
# ════════════════════════════════════════════════════════════════
# 积分相关路由 (prefix=/points)
# ════════════════════════════════════════════════════════════════
@points_router.get("/balance", response_model=PointsBalanceResponse)
def get_balance(
@router.get("/balance", response_model=PointsBalanceResponse)
async def get_balance(
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""查询当前用户积分余额 + 会员状态。"""
svc = _get_service()
account = svc.get_or_create_account(current_user.user.id, db)
svc: PointsService = Depends(get_points_service),
) -> PointsBalanceResponse:
user = current_user.user
account = svc.get_account(user.id)
limit = _member_limit(current_user)
if limit == -1:
usage = type("U", (), {"used": 0, "limit": -1, "remaining": -1, "reset_at": None})()
else:
usage = svc.get_daily_usage(user.id, limit=limit)
reset_at = usage.reset_at or datetime.combine(
datetime.now(CST).date() + timedelta(days=1),
datetime.min.time(),
tzinfo=CST,
)
return PointsBalanceResponse(
balance=account["balance"],
total_earned=account["total_earned"],
total_spent=account["total_spent"],
is_member=_is_member(current_user),
member_type=_member_type(current_user),
member_expires_at=getattr(current_user.user, "member_expires_at", None),
balance=int(account.balance or 0),
total_earned=int(account.total_earned or 0),
total_spent=int(account.total_spent or 0),
is_member=bool(getattr(user, "is_member", False)),
member_type=getattr(user, "member_type", None),
member_expires_at=getattr(user, "member_expires_at", None),
daily_free_clips_used=int(usage.used or 0),
daily_free_clips_limit=int(usage.limit),
daily_free_clips_remaining=int(usage.remaining if usage.limit != -1 else -1),
daily_reset_at=reset_at,
)
@points_router.get("/transactions", response_model=PointsTransactionsResponse)
def get_transactions(
page: int = Query(1, ge=1),
page_size: int = Query(20, ge=1, le=100),
type: Optional[str] = Query(None, description="筛选类型: add/deduct"),
source: Optional[str] = Query(None, description="筛选来源场景"),
start_date: Optional[datetime] = Query(None),
end_date: Optional[datetime] = Query(None),
# ── 流水 ──────────────────────────────────────────────────────────────────
@router.get("/transactions", response_model=PointsTransactionListResponse)
async def list_transactions(
page: int = Query(1, ge=1, description="页码"),
page_size: int = Query(20, ge=1, le=100, description="每页条数"),
type: Optional[str] = Query(None, description="流水类型: earn/spend/refund"),
source: Optional[str] = Query(None, description="流水来源/场景"),
start_date: Optional[datetime] = Query(None, description="起始时间(ISO8601)"),
end_date: Optional[datetime] = Query(None, description="结束时间(ISO8601)"),
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""查询积分流水(分页+筛选)。"""
svc = _get_service()
result = svc.get_transactions(
user_id=current_user.user.id,
db=db,
page=page,
page_size=page_size,
type_filter=type,
source_filter=source,
svc: PointsService = Depends(get_points_service),
) -> PointsTransactionListResponse:
offset = (page - 1) * page_size
items, total = svc.list_transactions(
current_user.user.id,
offset=offset,
limit=page_size,
type_=type,
source=source,
start_date=start_date,
end_date=end_date,
)
return PointsTransactionsResponse(**result)
return PointsTransactionListResponse(
items=[
PointsTransactionItem(
id=tx.id,
type=tx.type,
source=tx.source,
amount=int(tx.amount or 0),
balance_after=int(tx.balance_after or 0),
description=tx.description or "",
ref_id=tx.ref_id or "",
created_at=tx.created_at or datetime.now(timezone.utc),
)
for tx in items
],
total=int(total),
)
@points_router.get("/rules", response_model=PointsRulesResponse)
def get_rules(
_current_user: AuthenticatedUser = Depends(get_current_user),
):
"""查询所有积分消耗规则。"""
rules = []
for scene_key, scene_data in POINTS_SCENES.items():
# ── 积分包 ────────────────────────────────────────────────────────────────
_PACKAGE_DESC = {
"starter_pack": "新用户体验包,足够尝试多次智能混剪",
"basic_pack": "适合轻度使用,性价比之选",
"pro_pack": "重度创作者推荐,单积分单价更低",
}
@router.get("/packages", response_model=PointsPackagesResponse)
async def list_packages(
current_user: AuthenticatedUser = Depends(get_current_user),
svc: PointsService = Depends(get_points_service),
) -> PointsPackagesResponse:
user = current_user.user
member_type = getattr(user, "member_type", None) if bool(getattr(user, "is_member", False)) else None
discount = _member_discount(member_type)
from packages.domain.points import POINTS_PACKAGES, calc_package_price
pkgs: list[PointsPackageItem] = []
for pkg in POINTS_PACKAGES:
discounted_cents, original_cents, _ = calc_package_price(pkg["id"], member_type)
pkgs.append(
PointsPackageItem(
id=pkg["id"],
name=pkg["name"],
points=int(pkg["points"]),
price_cents=int(original_cents),
discounted_price_cents=int(discounted_cents),
currency="CNY",
description=_PACKAGE_DESC.get(pkg["id"], ""),
)
)
return PointsPackagesResponse(
packages=pkgs,
user_discount=float(discount),
unit_price_yuan=float(POINTS_UNIT_PRICE_YUAN),
)
# ── 充值下单 ──────────────────────────────────────────────────────────────
@router.post("/recharge", response_model=PointsRechargeResponse)
async def recharge(
req: PointsRechargeRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: PointsService = Depends(get_points_service),
) -> PointsRechargeResponse:
user = current_user.user
member_type = getattr(user, "member_type", None) if bool(getattr(user, "is_member", False)) else None
try:
order = svc.create_order(
user_id=user.id,
package_id=req.package_id,
payment_method=req.payment_method,
member_type_for_discount=member_type,
)
except ValueError as e:
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
return PointsRechargeResponse(
order_id=order.id,
package_name=order.package_name,
points_amount=int(order.points_amount or 0),
price_cents=int(order.price_cents or 0),
discount=float(order.discount or 1.0),
payment_params={}, # P5 接入微信支付后填充
)
# ── 扣减预估(不实际扣减) ───────────────────────────────────────────────
@router.post("/check", response_model=PointsDeductResponse)
async def check_deduct(
req: PointsDeductRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: PointsService = Depends(get_points_service),
) -> PointsDeductResponse:
user = current_user.user
is_member = bool(getattr(user, "is_member", False))
# 免费额度:ai_video 且未超限时视为 free_quota
is_free_quota = False
required = svc.calculate_cost(
req.scene_key,
is_member=is_member,
duration_minutes=req.duration_minutes,
)
if req.scene_key == "ai_video" and not is_member and required > 0:
usage = svc.get_daily_usage(user.id, limit=FREE_DAILY_CLIPS)
if int(usage.remaining or 0) > 0:
is_free_quota = True
balance = svc.get_balance(user.id)
# 免费额度下实际所需积分 = 0
real_required = 0 if is_free_quota else required
remaining_after = balance - real_required
allowed = is_free_quota or balance >= required
return PointsDeductResponse(
allowed=bool(allowed),
required_points=int(required),
current_balance=int(balance),
remaining_after=int(remaining_after if remaining_after >= 0 else balance),
transaction_id=None,
is_free_quota=bool(is_free_quota),
)
# ── 退还积分 ──────────────────────────────────────────────────────────────
@router.post("/refund", response_model=PointsRefundResponse)
async def refund(
req: PointsRefundRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: PointsService = Depends(get_points_service),
) -> PointsRefundResponse:
ok = svc.refund(current_user.user.id, req.transaction_id, reason=req.reason)
if ok:
try:
svc._accounts.session.commit() # type: ignore[attr-defined]
except Exception:
svc._accounts.session.rollback() # type: ignore[attr-defined]
raise
return PointsRefundResponse(success=bool(ok))
# ── 规则 ──────────────────────────────────────────────────────────────────
@router.get("/rules", response_model=PointsRulesResponse)
async def get_rules(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> PointsRulesResponse:
rules: list[PointsRuleItem] = []
for key, r in POINTS_RULES.items():
base = int(r.get("base_points", 0))
extra = r.get("extra_per_30s", 0) or 0
rules.append(
PointRuleItem(
scene_key=scene_key,
name=scene_data["name"],
base_points=scene_data["base_points"],
unit=scene_data["unit"],
extra_per_30s=scene_data.get("extra_per_30s"),
PointsRuleItem(
scene_key=key,
scene_name=r["name"],
points_per_use=base,
unit=r["unit"],
extra_per_30s=int(extra) if extra else None,
)
)
return PointsRulesResponse(
rules=rules,
free_user_multiplier=FREE_USER_MULTIPLIER,
free_user_multiplier=float(FREE_USER_MULTIPLIER),
note="免费用户消耗倍率为 {:.2f},付费会员按会员价扣减;ai_video 按每 30s 阶梯计费,含每日 {} 条免费额度。".format(
FREE_USER_MULTIPLIER, FREE_DAILY_CLIPS
),
)
@points_router.get("/packages", response_model=PointsPackagesResponse)
def get_packages(
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""查询可购买的积分包列表。"""
packages = []
for code, pkg in POINTS_PACKAGES.items():
unit_price = f"¥{pkg['price_cents'] / 100 / pkg['points']:.3f}/积分"
packages.append(
PointsPackageItem(
code=code,
name=pkg["name"],
points=pkg["points"],
price_cents=pkg["price_cents"],
unit_price=unit_price,
)
)
mt = _member_type(current_user)
discount = MEMBER_DISCOUNT.get(mt) if mt else None
return PointsPackagesResponse(packages=packages, user_discount=discount)
# ── 支付回调 stub ────────────────────────────────────────────────────────
@points_router.post("/check", response_model=PointsCheckResponse)
def check_points(
body: PointsCheckRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""消费前检查余额是否足够。"""
is_mem = _is_member(current_user)
mt = _member_type(current_user)
# 混剪场景先检查免费额度
is_free_quota = False
if body.scene_key == "ai_video" and not is_mem:
svc = _get_service()
if svc.check_daily_free_clip(current_user.user.id, db):
is_free_quota = True
required = calculate_points_cost(
body.scene_key,
is_mem,
quantity=body.quantity or 1,
duration_minutes=body.duration_minutes or 0,
member_type=mt,
)
svc = _get_service()
account = svc.get_or_create_account(current_user.user.id, db)
balance = account["balance"]
return PointsCheckResponse(
allowed=is_free_quota or balance >= required,
required_points=required,
current_balance=balance,
remaining_after=balance - required,
is_free_quota=is_free_quota,
)
@points_router.post("/deduct", response_model=SimpleMessageResponse)
def deduct_points(
body: PointsDeductRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""积分扣减(内部服务调用)。"""
svc = _get_service()
result = svc.deduct_points(
user_id=current_user.user.id,
amount=body.amount,
source=body.scene_key,
db=db,
description=body.description or "",
ref_id=body.ref_id or "",
)
if not result["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {body.amount},余额 {result['balance']}",
},
)
return SimpleMessageResponse(
success=True,
message=f"扣减 {body.amount} 积分成功",
data={"transaction_id": result["transaction_id"], "balance": result["balance"]},
)
@points_router.post("/refund", response_model=SimpleMessageResponse)
def refund_points(
body: PointsRefundRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""积分退还(内部服务调用)。"""
from packages.adapters.sqlalchemy_impl.models import PointsTransactionModel
txn = (
db.query(PointsTransactionModel)
.filter(PointsTransactionModel.id == body.transaction_id)
.first()
)
if txn is None:
raise HTTPException(status_code=404, detail="交易记录不存在")
if txn.user_id != current_user.user.id:
raise HTTPException(status_code=403, detail="无权退还他人积分")
svc = _get_service()
result = svc.refund_points(
user_id=current_user.user.id,
amount=txn.amount,
source=txn.source,
db=db,
ref_id=body.transaction_id,
description=body.reason or f"退还: {txn.description}",
)
if not result["success"]:
raise HTTPException(status_code=500, detail="退还失败")
return SimpleMessageResponse(
success=True,
message=f"退还 {txn.amount} 积分成功",
data={"transaction_id": result["transaction_id"], "balance": result["balance"]},
)
@points_router.post("/recharge", response_model=PointsOrderResponse)
def create_recharge_order(
body: PointsRechargeRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""创建积分充值订单。"""
svc = _get_service()
try:
order = svc.create_order(
user_id=current_user.user.id,
order_type="points",
product_code=body.package_id,
db=db,
)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) from None
return PointsOrderResponse(**order)
@points_router.get("/subscription/membership", response_model=MembershipStatusResponse)
def get_membership_status(
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""获取当前用户会员状态(聚合信息)。"""
svc = _get_service()
account = svc.get_or_create_account(current_user.user.id, db)
is_mem = _is_member(current_user)
max_resolution = "1080p" if is_mem else "720p"
return MembershipStatusResponse(
is_member=is_mem,
member_type=_member_type(current_user),
member_expires_at=getattr(current_user.user, "member_expires_at", None),
points_balance=account["balance"],
max_resolution=max_resolution,
)
# ════════════════════════════════════════════════════════════════
# 每日额度路由 (prefix=/usage)
# ════════════════════════════════════════════════════════════════
@usage_router.get("/daily", response_model=DailyUsageResponse)
def get_daily_usage(
current_user: AuthenticatedUser = Depends(get_current_user),
db: Session = Depends(get_db_session),
):
"""查询今日免费混剪额度使用情况。"""
svc = _get_service()
result = svc.get_daily_usage(current_user.user.id, db)
return DailyUsageResponse(**result)
# 为了向后兼容,也导出一个不带后缀的 router(方便旧引用)
router = points_router
@router.post("/payment-callback")
async def payment_callback_stub() -> dict:
"""微信支付回调占位(P5 接入真实签名校验与订单履约)。"""
return {"received": True}
+173 -146
View File
@@ -1,32 +1,45 @@
"""Subscription management API routes."""
"""会员订阅 API 路由(#1895 P3 简化版:免费 / 付费两档,付费分月/季/年)。
保留旧版 billing/payment-callback 相关端点和辅助函数以兼容现有单元测试。
新前端对接 /subscription/current、/plans、/subscribe、/cancel 即可。
"""
from __future__ import annotations
import logging
from dataclasses import replace
from datetime import datetime, timezone
from typing import List
import uuid
from datetime import datetime, timedelta, timezone
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_user_repository
from app.dependencies import get_points_service
from app.schemas.subscription import (
BillingRecord,
ChangePlanRequest,
ChangePlanResponse,
MembershipPlanItem,
MembershipPlansResponse,
SimpleResponse,
SubscriptionInfo,
SubscribeRequest,
SubscribeResponse,
SubscriptionInfoResponse,
ToggleAutoRenewRequest,
)
from fastapi import APIRouter, Depends, HTTPException, status
from packages.ports.user_repository import UserRepository
from packages.application.points_service import MEMBERSHIP_PLANS, PointsService
from packages.domain.points import FREE_DAILY_CLIPS
logger = logging.getLogger(__name__)
router = APIRouter()
MEMBER_TYPE_NAME = {
"monthly": "月度会员",
"quarterly": "季度会员",
"yearly": "年度会员",
}
# ── 旧版兼容:PLAN_QUOTAS / 套餐名 / 价格(保留以不破坏既有测试与旧前端) ──
# ============ 配额定义(硬编码,后续可迁移到配置中心) ============
PLAN_QUOTAS = {
"free": {"max_projects": 3, "max_storage_gb": 10},
@@ -36,11 +49,8 @@ PLAN_QUOTAS = {
}
# ============ Helper Functions ============
def _get_plan_name(plan_id: str) -> str:
"""获取套餐显示名称"""
"""获取套餐显示名称(旧版 4 档,保留兼容)。"""
plan_names = {
"free": "体验版",
"standard": "标准版",
@@ -51,7 +61,7 @@ def _get_plan_name(plan_id: str) -> str:
def _get_plan_price(plan_id: str, billing_cycle: str) -> float:
"""获取套餐价格"""
"""获取套餐价格(旧版 4 档,保留兼容)。"""
prices = {
("free", "monthly"): 0,
("free", "yearly"): 0,
@@ -65,145 +75,183 @@ def _get_plan_price(plan_id: str, billing_cycle: str) -> float:
return prices.get((plan_id, billing_cycle), 0)
def _build_subscription_info(user: AuthenticatedUser) -> SubscriptionInfo:
"""构建订阅信息响应"""
now = datetime.now(timezone.utc)
if user.user.subscription_expires_at:
period_end = user.user.subscription_expires_at.isoformat()
period_start = now.isoformat()
else:
period_start = now.isoformat()
period_end = now.isoformat()
def _member_type_name(member_type: str | None) -> str:
if not member_type:
return "免费会员"
return MEMBER_TYPE_NAME.get(member_type, "付费会员")
return SubscriptionInfo(
id=f"sub-{user.user.id[:8]}",
plan_id=user.user.subscription_plan or "free",
plan_name=_get_plan_name(user.user.subscription_plan or "free"),
status=user.user.subscription_status or "active",
billing_cycle="monthly",
current_period_start=period_start,
current_period_end=period_end,
amount=_get_plan_price(user.user.subscription_plan or "free", "monthly"),
auto_renew=True,
created_at=user.user.created_at.isoformat() if user.user.created_at else now.isoformat(),
def _effective_is_member(user) -> bool:
"""会员有效判定:is_member=True 且未过期。过期的视为免费用户。"""
if not bool(getattr(user, "is_member", False)):
return False
expires = getattr(user, "member_expires_at", None)
if expires is None:
return True
# 比较需带时区
now = datetime.now(timezone.utc)
if expires.tzinfo is None:
expires = expires.replace(tzinfo=timezone.utc)
return expires > now
# ── 当前会员状态 ─────────────────────────────────────────────────────────
@router.get("/current", response_model=SubscriptionInfoResponse)
async def get_current_subscription(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> SubscriptionInfoResponse:
user = current_user.user
is_member = _effective_is_member(user)
member_type = getattr(user, "member_type", None)
member_expires_at = getattr(user, "member_expires_at", None)
points_balance = int(getattr(user, "points_balance", 0) or 0)
daily_limit = -1 if is_member else FREE_DAILY_CLIPS
return SubscriptionInfoResponse(
is_member=is_member,
member_type=member_type if is_member else None,
member_type_name=_member_type_name(member_type if is_member else None),
member_expires_at=member_expires_at if is_member else None,
points_balance=points_balance,
daily_free_clips_limit=daily_limit,
)
# ============ API Endpoints ============
# ── 会员套餐列表 ─────────────────────────────────────────────────────────
@router.get("/current", response_model=SubscriptionInfo)
async def get_current_subscription(
@router.get("/plans", response_model=MembershipPlansResponse)
async def list_plans(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> SubscriptionInfo:
"""获取当前订阅信息"""
return _build_subscription_info(current_user)
) -> MembershipPlansResponse:
_plan_desc = {
"monthly": "适合轻度创作者,30+1 天会员期",
"quarterly": "高性价比之选,93 天会员期",
"yearly": "重度创作者推荐,366 天会员期,享最高折扣",
}
plans: list[MembershipPlanItem] = []
for key, plan in MEMBERSHIP_PLANS.items():
plans.append(
MembershipPlanItem(
member_type=key,
name=plan["name"],
price_cents=int(plan["price_cents"]),
days=int(plan["days"]),
discount=float(plan["discount"]),
daily_free_clips_limit=int(plan["daily_free_clips_limit"]),
description=_plan_desc.get(key, ""),
)
)
return MembershipPlansResponse(plans=plans)
@router.get("/billing-records", response_model=List[BillingRecord])
# ── 订阅下单 ─────────────────────────────────────────────────────────────
@router.post("/subscribe", response_model=SubscribeResponse)
async def subscribe(
req: SubscribeRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
svc: PointsService = Depends(get_points_service),
) -> SubscribeResponse:
if req.member_type not in MEMBERSHIP_PLANS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"无效的会员类型,支持: {', '.join(MEMBERSHIP_PLANS.keys())}",
)
order = svc.create_membership_order(
user_id=current_user.user.id,
member_type=req.member_type,
payment_method=req.payment_method,
)
plan = MEMBERSHIP_PLANS[req.member_type]
return SubscribeResponse(
order_id=order.id,
member_type=req.member_type,
member_type_name=_member_type_name(req.member_type),
price_cents=int(plan["price_cents"]),
discount=float(plan["discount"]),
payment_params={}, # P5 接入微信支付
)
# ── 取消自动续费(占位) ─────────────────────────────────────────────────
@router.post("/cancel", response_model=SimpleResponse)
async def cancel_auto_renew(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> SimpleResponse:
"""取消自动续费(占位:首期不做自动续费,直接返回成功)。"""
return SimpleResponse(success=True, message="已取消自动续费")
# ── 旧版端点兼容 ─────────────────────────────────────────────────────────
# 以下端点保留旧路径与返回结构以兼容旧前端和现有单元测试;
# 新逻辑走 /subscribe + /payment-callback(stub)。
@router.get("/billing-records")
async def get_billing_records(
current_user: AuthenticatedUser = Depends(get_current_user),
) -> List[BillingRecord]:
"""获取账单记录列表"""
) -> list[dict]:
"""获取账单记录列表(旧版端点,暂时返回空列表)。"""
from packages.adapters.sqlalchemy_impl.billing_repository import SQLAlchemyBillingRepository
from packages.adapters.sqlalchemy_impl.session import SessionLocal
if SessionLocal is None:
return []
session = SessionLocal()
try:
session = SessionLocal()
except Exception:
return []
try:
repo = SQLAlchemyBillingRepository(session)
records = repo.find_by_user(current_user.user.id)
return [
BillingRecord(
id=r.id,
plan_name=r.plan_name,
amount=r.amount,
billing_cycle=r.billing_cycle,
status=r.status,
payment_method=r.payment_method or "未支付",
created_at=r.created_at.isoformat() if r.created_at else "",
invoice_url=r.invoice_url,
)
{
"id": r.id,
"plan_name": r.plan_name,
"amount": r.amount,
"billing_cycle": r.billing_cycle,
"status": r.status,
"payment_method": r.payment_method or "未支付",
"created_at": r.created_at.isoformat() if r.created_at else "",
"invoice_url": getattr(r, "invoice_url", None),
}
for r in records
]
finally:
session.close()
@router.post("/change-plan", response_model=ChangePlanResponse)
@router.post("/change-plan")
async def change_plan(
request: ChangePlanRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> ChangePlanResponse:
"""变更订阅套餐(升级/降级)"""
# TODO: 接入支付验证(支付宝/微信支付)
valid_plans = {"free", "standard", "pro", "enterprise"}
) -> dict:
"""变更套餐(旧版端点,提示新版走 /subscribe)。"""
valid_plans = {"free", "standard", "pro", "enterprise", "paid"}
if request.target_plan_id not in valid_plans:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"无效的套餐ID。支持的套餐: {', '.join(valid_plans)}",
status_code=400,
detail=f"无效的套餐ID。支持的套餐: {', '.join(sorted(valid_plans))}",
)
valid_cycles = {"monthly", "yearly"}
if request.billing_cycle not in valid_cycles:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="无效的计费周期。支持: monthly, yearly",
)
user = current_user.user
current_plan = user.subscription_plan or "free"
target_plan = request.target_plan_id
if current_plan == target_plan:
return ChangePlanResponse(
success=False,
message=f"您已经是 {_get_plan_name(target_plan)}",
)
# 通过 dataclasses.replace 创建新实例(不直接修改 dataclass)
quotas = PLAN_QUOTAS.get(target_plan, PLAN_QUOTAS["free"])
updated_user = replace(
user,
subscription_plan=target_plan,
subscription_status="active",
max_projects=quotas["max_projects"],
max_storage_gb=quotas["max_storage_gb"],
)
user_repository.save(updated_user)
# 用更新后的用户构造响应
refreshed_auth_user = AuthenticatedUser(user=updated_user)
return ChangePlanResponse(
success=True,
message=f"套餐已成功变更为 {_get_plan_name(target_plan)}",
new_subscription=_build_subscription_info(refreshed_auth_user),
)
return {
"success": False,
"message": "旧版套餐已下线,请使用 /subscription/subscribe 订阅会员",
}
@router.post("/cancel", response_model=SimpleResponse)
async def cancel_subscription(
@router.post("/toggle-auto-renew")
async def toggle_auto_renew(
request: ToggleAutoRenewRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
user_repository: UserRepository = Depends(get_user_repository),
) -> SimpleResponse:
"""取消订阅"""
user = current_user.user
if user.subscription_plan == "free":
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail="体验版无需取消",
)
updated_user = replace(user, subscription_status="cancelled")
user_repository.save(updated_user)
return SimpleResponse(
success=True,
message="订阅已取消,当前周期结束后停止服务",
message="已开启自动续费" if request.enabled else "已关闭自动续费",
)
@@ -216,24 +264,23 @@ async def payment_callback(
payment_method: str = "alipay",
payment_id: str = "",
) -> dict:
"""支付回调 - 在事务中更新账单和订阅状态
"""支付回调(旧版端点保留,走 BillingRepository 旧逻辑)。
注意:生产环境需要验证支付签名
注意:新版会员订阅回调走 points service 的 mark_order_paid。
生产环境需要验证支付签名。
"""
import uuid
from datetime import timedelta
from packages.adapters.sqlalchemy_impl.billing_repository import SQLAlchemyBillingRepository
from packages.adapters.sqlalchemy_impl.session import SessionLocal
if SessionLocal is None:
raise HTTPException(status_code=500, detail="Database not available")
session = SessionLocal()
try:
session = SessionLocal()
except Exception as e:
raise HTTPException(status_code=500, detail="Database not available") from e
try:
repo = SQLAlchemyBillingRepository(session)
# 创建账单记录
record_id = uuid.uuid4().hex
repo.create(
{
@@ -245,35 +292,15 @@ async def payment_callback(
"status": "pending",
}
)
# 在事务中标记支付成功并更新订阅
repo.mark_paid(record_id, payment_method, payment_id)
# 计算到期时间
days = 365 if billing_cycle == "yearly" else 30
expires_at = datetime.now(timezone.utc) + timedelta(days=days)
repo.update_subscription_on_payment(user_id, plan, expires_at)
session.commit()
return {"success": True, "message": "支付成功", "record_id": record_id}
except Exception as e:
session.rollback()
logger.error(f"支付回调处理失败: user_id={user_id}, plan={plan}, error={e}")
# 不返回原始异常信息,避免泄漏内部实现细节
logger.error("支付回调处理失败: user_id=%s, plan=%s, error=%s", user_id, plan, e)
raise HTTPException(status_code=500, detail="支付处理失败,请稍后重试") from e
finally:
session.close()
@router.post("/toggle-auto-renew", response_model=SimpleResponse)
async def toggle_auto_renew(
request: ToggleAutoRenewRequest,
current_user: AuthenticatedUser = Depends(get_current_user),
) -> SimpleResponse:
"""切换自动续费"""
# TODO: 实际需要在数据库中存储 auto_renew 字段
status_text = "已开启自动续费" if request.enabled else "已关闭自动续费"
return SimpleResponse(
success=True,
message=status_text,
)
+80
View File
@@ -10,6 +10,7 @@ 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 (
@@ -18,6 +19,7 @@ from app.dependencies import (
get_audio_url_signer,
get_cosyvoice_service,
get_db_session,
get_points_service,
get_project_repository,
get_voice_clone_profile_repository,
)
@@ -40,6 +42,7 @@ from packages.adapters.sqlalchemy_impl.tts_job_repository import (
SQLAlchemyTTSJobRepository,
)
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
from packages.application.points_service import PointsService
from packages.application.tts_job.streaming_service import TTSStreamingService
from packages.application.tts_job.use_cases import (
CreateTTSJobUseCase,
@@ -52,6 +55,7 @@ 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.voice_presets import list_voices
from packages.middleware.points_gate import _is_active_member
from packages.ports.asset_library_repository import AssetLibraryRepository
from packages.ports.asset_repository import AssetRepository
from packages.ports.project_repository import ProjectRepository
@@ -131,6 +135,7 @@ def synthesize(
repository: SQLAlchemyTTSJobRepository = Depends(_get_repository),
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
voice_clone_repo=Depends(get_voice_clone_profile_repository),
points_svc: PointsService = Depends(get_points_service),
) -> TTSSynthesizeResponse:
"""发起 TTS 合成任务。
@@ -192,6 +197,25 @@ def synthesize(
metadata=synthesis_meta,
)
# 积分扣费(#1895 P4):文本长度粗估时长,后续失败分支手动退款
_pts_tx_id: str | None = None
if settings.POINTS_ENABLED:
text_len = len(request.text or "")
duration_minutes = max(1.0, text_len / 240.0)
_is_m = _is_active_member(authenticated_user.user)
_dr = points_svc.check_and_deduct(
user_id=user_id,
scene_key="ai_voice",
duration_minutes=duration_minutes,
description="AI 配音合成",
is_member=_is_m,
)
if not _dr.success:
from packages.middleware.points_gate import _insufficient_points
raise _insufficient_points(_dr.amount, _dr.balance, "AI 配音")
_pts_tx_id = _dr.transaction_id
# 提交 CosyVoice 合成任务
workflow = TTSWorkflowService(
repository=repository,
@@ -205,6 +229,12 @@ def synthesize(
# 但 DB 异常、网络异常等意外错误可能逃逸。
# 与音色克隆接口保持一致:标记 failed,返回 201,不抛 500。
logger.error(f"TTS 合成异常: job_id={job.id}, error={e}", exc_info=True)
if _pts_tx_id:
try:
points_svc.refund(user_id, _pts_tx_id, reason="TTS 合成失败")
_pts_tx_id = None
except Exception:
logger.exception("TTS 合成失败后退款失败 user=%s tx=%s", user_id, _pts_tx_id)
try:
job = workflow.process_synthesis_failure(job.id, str(e))
except Exception as inner_e:
@@ -223,11 +253,24 @@ def synthesize(
celery_app.send_task("worker.process_tts_synthesis", args=[job.id])
except Exception as e:
# Celery 调度失败,标记 job 为 failed
if _pts_tx_id:
try:
points_svc.refund(user_id, _pts_tx_id, reason="TTS Celery 调度失败")
_pts_tx_id = None
except Exception:
logger.exception("Celery 调度失败后退款失败 user=%s tx=%s", user_id, _pts_tx_id)
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}")
# 若最终 job 为 failed 状态(前面两个兜底路径之一),退积分
if _pts_tx_id and getattr(job.status, "value", str(job.status)) == "failed":
try:
points_svc.refund(user_id, _pts_tx_id, reason="TTS 任务创建后立即失败")
except Exception:
logger.exception("TTS 失败状态退款失败 user=%s tx=%s", user_id, _pts_tx_id)
return TTSSynthesizeResponse(
job_id=job.id,
status=job.status,
@@ -555,6 +598,7 @@ def preview_tts(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
cosyvoice_service: CosyVoiceService = Depends(get_cosyvoice_service),
voice_clone_repo=Depends(get_voice_clone_profile_repository),
points_svc: PointsService = Depends(get_points_service),
) -> TTSPreviewResponse:
"""TTS 预览(试听)——同步合成,立即返回音频 URL。
@@ -578,6 +622,25 @@ def preview_tts(
)
actual_voice_id = profile.voice_id
# 积分扣费(#1895 P4):preview 短文本默认 1 分钟
_pts_tx_id: str | None = None
if settings.POINTS_ENABLED:
text_len = len(request.text or "")
duration_minutes = max(1.0, text_len / 240.0)
_is_m = _is_active_member(authenticated_user.user)
_dr = points_svc.check_and_deduct(
user_id=authenticated_user.user.id,
scene_key="ai_voice",
duration_minutes=duration_minutes,
description="AI 配音试听",
is_member=_is_m,
)
if not _dr.success:
from packages.middleware.points_gate import _insufficient_points
raise _insufficient_points(_dr.amount, _dr.balance, "AI 配音试听")
_pts_tx_id = _dr.transaction_id
try:
result = cosyvoice_service.synthesize_speech(
text=request.text,
@@ -587,15 +650,32 @@ def preview_tts(
language=getattr(request, "language", "zh-CN"),
)
except CosyVoiceError as e:
if _pts_tx_id:
try:
points_svc.refund(authenticated_user.user.id, _pts_tx_id, reason="TTS 试听 CosyVoice 失败")
except Exception:
logger.exception("TTS 试听失败退款失败")
raise HTTPException(
status_code=status.HTTP_502_BAD_GATEWAY,
detail=f"TTS 合成失败: {e}",
) from e
except ValueError as e:
if _pts_tx_id:
try:
points_svc.refund(authenticated_user.user.id, _pts_tx_id, reason="TTS 试听参数错误")
except Exception:
logger.exception("TTS 试听失败退款失败")
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=str(e),
) from e
except Exception:
if _pts_tx_id:
try:
points_svc.refund(authenticated_user.user.id, _pts_tx_id, reason="TTS 试听异常")
except Exception:
logger.exception("TTS 试听失败退款失败")
raise
return TTSPreviewResponse(
audio_url=result.audio_url,
+53
View File
@@ -0,0 +1,53 @@
"""使用量 API 路由(#1895 P3) — 每日免费混剪额度等。"""
from __future__ import annotations
import logging
from datetime import datetime, timedelta, timezone
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_points_service
from app.schemas.points import DailyUsageResponse
from fastapi import APIRouter, Depends
from packages.application.points_service import FREE_DAILY_CLIPS, PointsService
logger = logging.getLogger(__name__)
router = APIRouter()
CST = timezone(timedelta(hours=8))
@router.get("/daily", response_model=DailyUsageResponse)
async def get_daily_usage(
current_user: AuthenticatedUser = Depends(get_current_user),
svc: PointsService = Depends(get_points_service),
) -> DailyUsageResponse:
"""获取当前用户今日免费混剪额度使用情况。
付费会员 daily_free_clips_limit = -1(不限),免费用户按 FREE_DAILY_CLIPS。
"""
user = current_user.user
is_member = bool(getattr(user, "is_member", False))
if is_member:
# 不限
today = datetime.now(CST).date()
reset_at = datetime.combine(today + timedelta(days=1), datetime.min.time(), tzinfo=CST)
return DailyUsageResponse(
free_clips_used=0,
free_clips_limit=-1,
free_clips_remaining=-1,
reset_at=reset_at,
)
usage = svc.get_daily_usage(user.id, limit=FREE_DAILY_CLIPS)
reset_at = usage.reset_at
if reset_at is None:
today = datetime.now(CST).date()
reset_at = datetime.combine(today + timedelta(days=1), datetime.min.time(), tzinfo=CST)
return DailyUsageResponse(
free_clips_used=int(usage.used or 0),
free_clips_limit=int(usage.limit),
free_clips_remaining=int(usage.remaining or 0),
reset_at=reset_at,
)
+18
View File
@@ -6,11 +6,13 @@ import logging
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_points_service,
get_project_repository,
get_voice_clone_profile_repository,
)
@@ -27,6 +29,7 @@ from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
SQLAlchemyVoiceCloneProfileRepository,
)
from packages.application.cosyvoice_service import CosyVoiceError, CosyVoiceService
from packages.application.points_service import PointsService
from packages.application.voice_clone.use_cases import (
DeleteVoiceCloneUseCase,
GetVoiceCloneStatusUseCase,
@@ -38,6 +41,7 @@ from packages.application.voice_clone.use_cases import (
from packages.application.voice_clone.workflow import (
VoiceCloneWorkflowService,
)
from packages.middleware.points_gate import _is_active_member
from packages.ports.asset_repository import AssetRepository
from packages.ports.project_repository import ProjectRepository
from packages.shared.storage import SharedStorageService
@@ -95,6 +99,7 @@ def create_voice_clone(
asset_repository: AssetRepository = Depends(get_asset_repository),
project_repository: ProjectRepository = Depends(get_project_repository),
storage_service: SharedStorageService = Depends(get_storage_service),
points_svc: PointsService = Depends(get_points_service),
) -> VoiceCloneProfileResponse:
"""创建音色克隆任务。
@@ -107,6 +112,19 @@ def create_voice_clone(
"""
user_id = authenticated_user.user.id
# 声音克隆训练免费(base_points=0),仅做记录打点;不退款
if settings.POINTS_ENABLED:
try:
points_svc.check_and_deduct(
user_id=user_id,
scene_key="voice_clone_train",
duration_minutes=1,
description="声音克隆训练",
is_member=_is_active_member(authenticated_user.user),
)
except Exception:
logger.exception("voice_clone_train 打点失败(不阻断业务)")
source_audio_url = request.source_audio_url
clone_metadata = dict(request.metadata_ or {})
+38
View File
@@ -25,6 +25,9 @@ from packages.adapters.sqlalchemy_impl.classification_job_repository import (
from packages.adapters.sqlalchemy_impl.cover_template_repository import (
SQLAlchemyCoverTemplateRepository,
)
from packages.adapters.sqlalchemy_impl.daily_usage_repository import (
SQLAlchemyDailyUsageRepository,
)
from packages.adapters.sqlalchemy_impl.duplication_repository import (
SQLAlchemyDuplicationRecordRepository,
)
@@ -38,6 +41,15 @@ from packages.adapters.sqlalchemy_impl.ingest_job_repository import (
SQLAlchemyIngestJobRepository,
)
from packages.adapters.sqlalchemy_impl.job_repository import SQLAlchemyJobRepository
from packages.adapters.sqlalchemy_impl.points_account_repository import (
SQLAlchemyPointsAccountRepository,
)
from packages.adapters.sqlalchemy_impl.points_order_repository import (
SQLAlchemyPointsOrderRepository,
)
from packages.adapters.sqlalchemy_impl.points_transaction_repository import (
SQLAlchemyPointsTransactionRepository,
)
from packages.adapters.sqlalchemy_impl.project_repository import (
SQLAlchemyProjectRepository,
)
@@ -53,6 +65,7 @@ from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
from packages.adapters.sqlalchemy_impl.voice_library_repository import (
SQLAlchemyVoiceLibraryRepository,
)
from packages.application.points_service import PointsService
from packages.ports.tag_repository import TagRepository
from packages.ports.user_repository import UserRepository
@@ -233,3 +246,28 @@ def get_audio_url_signer():
return storage.get_download_url(url, expires_seconds=86400)
return sign_audio_url
# ── Points / Membership 依赖 ──────────────────────────────────────────────
def get_redis_client():
"""提供通用 Redis 客户端(decode_responses=True),供 PointsService 等使用。"""
return redis.from_url(settings.REDIS_URL, decode_responses=True)
def _get_points_service_from_session(session: Session) -> PointsService:
return PointsService(
account_repo=SQLAlchemyPointsAccountRepository(session),
tx_repo=SQLAlchemyPointsTransactionRepository(session),
order_repo=SQLAlchemyPointsOrderRepository(session),
daily_usage_repo=SQLAlchemyDailyUsageRepository(session),
redis_client=redis.from_url(settings.REDIS_URL, decode_responses=True),
)
def get_points_service(
session: Session = Depends(get_db_session),
) -> PointsService:
"""提供 PointsService 实例(四个 points 仓储共享同一 DB session)。"""
return _get_points_service_from_session(session)
+94 -140
View File
@@ -1,182 +1,136 @@
"""积分 & 会员相关 Pydantic Schema (#1895)"""
"""积分/会员 P3 API 请求/响应 Schema."""
from __future__ import annotations
from datetime import datetime
from typing import Any, Optional
from typing import Optional
from pydantic import BaseModel, Field
# ============ 余额 & 账户 ============
# ── 余额 & 会员状态 ────────────────────────────────────────────────────────
class PointsBalanceResponse(BaseModel):
"""积分余额 + 会员状态"""
balance: int = Field(..., description="当前积分余额")
total_earned: int = Field(..., description="累计获得积分")
total_spent: int = Field(..., description="累计消耗积分")
is_member: bool = Field(default=False, description="是否付费会员")
member_type: Optional[str] = Field(None, description="会员类型: monthly/quarterly/yearly")
member_expires_at: Optional[datetime] = Field(None, description="会员到期时间")
balance: int
total_earned: int
total_spent: int
is_member: bool
member_type: Optional[str] = None
member_expires_at: Optional[datetime] = None
daily_free_clips_used: int
daily_free_clips_limit: int
daily_free_clips_remaining: int
daily_reset_at: datetime
# ============ 流水 ============
# ── 流水 ──────────────────────────────────────────────────────────────────
class PointsTransactionItem(BaseModel):
"""单条积分流水"""
id: str
type: str = Field(..., description="类型: add/deduct")
source: str = Field(..., description="来源场景")
type: str # earn / spend / refund
source: str
amount: int
balance_after: int
description: str = ""
ref_id: str = ""
created_at: Optional[str] = None
created_at: datetime
class PointsTransactionsResponse(BaseModel):
"""积分流水分页响应"""
class PointsTransactionListResponse(BaseModel):
items: list[PointsTransactionItem]
total: int
page: int
page_size: int
# ============ 规则 & 积分包 ============
# ── 积分包 ────────────────────────────────────────────────────────────────
class PointRuleItem(BaseModel):
"""单条积分规则"""
scene_key: str
class PointsPackageItem(BaseModel):
id: str
name: str
base_points: int
points: int
price_cents: int
discounted_price_cents: int = 0
currency: str = "CNY"
description: str = ""
class PointsPackagesResponse(BaseModel):
packages: list[PointsPackageItem]
user_discount: float = 1.0
unit_price_yuan: float = 0.10
# ── 充值 ──────────────────────────────────────────────────────────────────
class PointsRechargeRequest(BaseModel):
package_id: str = Field(..., description="积分包ID: starter_pack/basic_pack/pro_pack")
payment_method: str = Field("wechat_pay", description="支付方式,预留 wechat_pay")
class PointsRechargeResponse(BaseModel):
order_id: str
package_name: str
points_amount: int
price_cents: int
discount: float
payment_params: dict = Field(default_factory=dict)
# ── 扣减预估(内部) ──────────────────────────────────────────────────────
class PointsDeductRequest(BaseModel):
scene_key: str
duration_minutes: float = 1.0
ref_id: str = ""
description: str = ""
class PointsDeductResponse(BaseModel):
allowed: bool
required_points: int
current_balance: int
remaining_after: int
transaction_id: Optional[str] = None
is_free_quota: bool = False
# ── 退还 ──────────────────────────────────────────────────────────────────
class PointsRefundRequest(BaseModel):
transaction_id: str
reason: str = ""
class PointsRefundResponse(BaseModel):
success: bool
# ── 规则 ──────────────────────────────────────────────────────────────────
class PointsRuleItem(BaseModel):
scene_key: str
scene_name: str
points_per_use: int
unit: str
extra_per_30s: Optional[int] = None
class PointsRulesResponse(BaseModel):
"""所有积分消耗规则"""
rules: list[PointRuleItem]
free_user_multiplier: float = Field(..., description="免费用户积分上浮系数")
rules: list[PointsRuleItem]
free_user_multiplier: float
note: str = ""
class PointsPackageItem(BaseModel):
"""积分包信息"""
code: str
name: str
points: int
price_cents: int
unit_price: str = Field("", description="单价描述,如 ¥0.099/积分")
class PointsPackagesResponse(BaseModel):
"""可购买的积分包列表"""
packages: list[PointsPackageItem]
user_discount: Optional[float] = Field(None, description="当前用户折扣(会员)")
# ============ 消费前检查 ============
class PointsCheckRequest(BaseModel):
"""消费前余额检查请求"""
scene_key: str
duration_minutes: Optional[float] = None
quantity: Optional[int] = 1
class PointsCheckResponse(BaseModel):
"""消费前余额检查响应"""
allowed: bool
required_points: int
current_balance: int
remaining_after: int
is_free_quota: bool = False
# ============ 手动扣减 / 退还(内部接口) ============
class PointsDeductRequest(BaseModel):
"""积分扣减请求"""
scene_key: str
amount: int
description: Optional[str] = ""
ref_id: Optional[str] = ""
class PointsRefundRequest(BaseModel):
"""积分退还请求"""
transaction_id: str
reason: Optional[str] = ""
class PointsRechargeRequest(BaseModel):
"""积分充值请求"""
package_id: str = Field(..., description="积分包 code,如 starter_pack")
# ============ 订单 ============
class PointsOrderResponse(BaseModel):
"""订单信息"""
id: str
order_type: str
product_code: str
amount_cents: int
status: str
created_at: Optional[str] = None
# ============ 每日额度 ============
# ── 每日免费额度 ──────────────────────────────────────────────────────────
class DailyUsageResponse(BaseModel):
"""今日免费额度使用情况"""
free_clips_used: int
free_clips_limit: int
free_clips_remaining: int
reset_at: str
# ============ 会员状态(聚合) ============
class MembershipStatusResponse(BaseModel):
"""当前用户会员状态(聚合信息)"""
is_member: bool
member_type: Optional[str] = None
member_expires_at: Optional[datetime] = None
points_balance: int
max_resolution: str = Field(
default="1080p",
description="可用最高分辨率: 720p(free) / 1080p(paid)",
)
# ============ 通用响应 ============
class SimpleMessageResponse(BaseModel):
"""简单消息响应"""
success: bool
message: str
data: Optional[dict[str, Any]] = None
reset_at: datetime
+107 -82
View File
@@ -1,105 +1,130 @@
"""Subscription schemas for API request/response models."""
"""Subscription schemas(#1895 P3 简化为两档会员:免费 / 付费)."""
from __future__ import annotations
from datetime import datetime
from typing import List as _List
from typing import Optional
from pydantic import BaseModel, Field
# ============ Enums / Types ============
# ── 当前会员状态 ──────────────────────────────────────────────────────────
class PlanType(str):
"""套餐类型"""
FREE = "free"
STANDARD = "standard"
PRO = "pro"
ENTERPRISE = "enterprise"
class SubscriptionStatus(str):
"""订阅状态"""
ACTIVE = "active"
EXPIRED = "expired"
CANCELLED = "cancelled"
TRIAL = "trial"
class BillingStatus(str):
"""账单状态"""
PAID = "paid"
PENDING = "pending"
FAILED = "failed"
REFUNDED = "refunded"
class BillingCycle(str):
"""计费周期"""
MONTHLY = "monthly"
YEARLY = "yearly"
# ============ Response Schemas ============
class SubscriptionInfo(BaseModel):
"""当前订阅信息"""
id: str
plan_id: str
plan_name: str
status: str
billing_cycle: str
current_period_start: str
current_period_end: str
amount: float
auto_renew: bool
created_at: str
class BillingRecord(BaseModel):
"""账单记录"""
id: str
plan_name: str
amount: float
billing_cycle: str
status: str
payment_method: str
created_at: str
invoice_url: Optional[str] = None
class ChangePlanResponse(BaseModel):
"""升级/降级响应"""
success: bool
message: str
new_subscription: Optional[SubscriptionInfo] = None
class SubscriptionInfoResponse(BaseModel):
is_member: bool
member_type: Optional[str] = None # monthly / quarterly / yearly
member_type_name: str = "免费会员"
member_expires_at: Optional[datetime] = None
points_balance: int = 0
daily_free_clips_limit: int = 2 # -1 表示不限
class SimpleResponse(BaseModel):
"""简单响应(用于取消订阅、切换自动续费等)"""
"""通用简单成功响应."""
success: bool
message: str
message: str = ""
# ============ Request Schemas ============
# ── 订阅 ──────────────────────────────────────────────────────────────────
class ChangePlanRequest(BaseModel):
"""升级/降级请求"""
class SubscribeRequest(BaseModel):
member_type: str = Field(..., description="会员类型: monthly / quarterly / yearly")
payment_method: str = Field("wechat_pay", description="支付方式(预留)")
target_plan_id: str = Field(..., description="目标套餐ID")
billing_cycle: str = Field(..., description="计费周期: monthly/yearly")
class SubscribeResponse(BaseModel):
order_id: str
member_type: str
member_type_name: str
price_cents: int
discount: float
payment_params: dict = Field(default_factory=dict)
# ── 套餐列表 ──────────────────────────────────────────────────────────────
class MembershipPlanItem(BaseModel):
member_type: str
name: str
price_cents: int
days: int
discount: float
daily_free_clips_limit: int # -1 不限
description: str = ""
class MembershipPlansResponse(BaseModel):
plans: list[MembershipPlanItem]
# 保留旧的名称别名,兼容其它模块导入(内部不使用旧的 4 档枚举)
class ChangePlanResponse(SimpleResponse):
new_subscription: Optional["SubscriptionInfo"] = None
class ToggleAutoRenewRequest(BaseModel):
"""切换自动续费请求"""
enabled: bool
enabled: bool = Field(..., description="是否开启自动续费")
# ── 旧版 4 档 Schema 兼容别名(供集成测试 fixture 和历史代码导入) ────────
class SubscriptionInfo(BaseModel):
"""旧版订阅信息(集成测试 fixture 使用,新代码请用 SubscriptionInfoResponse)。"""
id: str = ""
plan_id: str = "free"
plan: str = "free"
plan_name: str = "体验版"
status: str = "active"
billing_cycle: str = "monthly"
current_period_start: str = ""
current_period_end: str = ""
amount: float = 0
amount_cents: int = 0
auto_renew: bool = True
created_at: str = ""
expires_at: Optional[datetime] = None
max_projects: int = -1
max_storage_gb: int = -1
is_member: bool = False
member_type: Optional[str] = None
member_expires_at: Optional[datetime] = None
points_balance: int = 0
daily_free_clips_limit: int = 2
class BillingRecord(BaseModel):
"""旧版账单记录(集成测试 fixture 使用)。"""
id: str = ""
plan_id: str = ""
plan_name: str = ""
billing_cycle: str = ""
amount: float = 0
amount_cents: int = 0
description: str = ""
status: str = "paid"
paid_at: Optional[str] = None
period_start: str = ""
period_end: str = ""
created_at: Optional[datetime] = None
class ChangePlanRequest(BaseModel):
"""旧版套餐变更请求(集成测试 fixture 使用)。"""
target_plan_id: str
billing_cycle: str
plan: str = "free"
class BillingRecordsResponse(BaseModel):
records: _List[BillingRecord] = []
ChangePlanResponse.model_rebuild()
@@ -1,32 +1,22 @@
from __future__ import annotations
from datetime import date, datetime, timezone
from uuid import uuid4
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import DailyUsageRecordModel
from packages.domain.daily_usage_record import DailyUsageRecord
class SQLAlchemyDailyUsageRepository:
"""每日使用计数仓储 — DB 持久化兜底;Redis 为实时计数主存储."""
def __init__(self, session: Session):
self.session = session
def create(self, record: DailyUsageRecord) -> DailyUsageRecord:
model = DailyUsageRecordModel(
id=record.id,
user_id=record.user_id,
usage_date=record.usage_date,
usage_type=record.usage_type,
count=record.count,
updated_at=record.updated_at,
)
self.session.add(model)
self.session.commit()
return record
def get_by_user_and_date(
self, user_id: str, usage_date: date, usage_type: str = "free_clip"
) -> DailyUsageRecord | None:
model = (
def get_for_today(self, user_id: str, usage_date: date, usage_type: str) -> DailyUsageRecordModel:
"""按 (user_id, date, type) 获取记录,不存在则 UPSERT 一条 count=0 的记录并返回."""
record = (
self.session.query(DailyUsageRecordModel)
.filter(
DailyUsageRecordModel.user_id == user_id,
@@ -35,59 +25,59 @@ class SQLAlchemyDailyUsageRepository:
)
.first()
)
if model is None:
return None
return self._to_domain(model)
def update_count(self, record: DailyUsageRecord) -> DailyUsageRecord:
model = self.session.query(DailyUsageRecordModel).filter(DailyUsageRecordModel.id == record.id).first()
if model is None:
if record is not None:
return record
model.count = record.count
model.updated_at = datetime.now(timezone.utc)
self.session.add(model)
self.session.commit()
record = DailyUsageRecordModel(
id=str(uuid4()),
user_id=user_id,
usage_date=usage_date,
usage_type=usage_type,
count=0,
updated_at=datetime.now(timezone.utc),
)
self.session.add(record)
try:
self.session.flush()
except Exception:
# 并发 UPSERT 冲突,回退到查询
self.session.rollback()
record = (
self.session.query(DailyUsageRecordModel)
.filter(
DailyUsageRecordModel.user_id == user_id,
DailyUsageRecordModel.usage_date == usage_date,
DailyUsageRecordModel.usage_type == usage_type,
)
.first()
)
if record is not None:
return record
raise
return record
def upsert(self, user_id: str, usage_date: date, usage_type: str = "free_clip") -> DailyUsageRecord:
"""Increment usage count for the given user/date/type, creating if needed."""
model = (
self.session.query(DailyUsageRecordModel)
.filter(
DailyUsageRecordModel.user_id == user_id,
DailyUsageRecordModel.usage_date == usage_date,
DailyUsageRecordModel.usage_type == usage_type,
)
.first()
)
if model is None:
record = DailyUsageRecord.create(user_id=user_id, usage_date=usage_date, usage_type=usage_type)
record.count = 1
model = DailyUsageRecordModel(
id=record.id,
user_id=record.user_id,
usage_date=record.usage_date,
usage_type=record.usage_type,
count=1,
updated_at=datetime.now(timezone.utc),
)
self.session.add(model)
self.session.commit()
return record
def increment_count(self, user_id: str, usage_date: date, usage_type: str, delta: int = 1) -> int:
"""原子地 count += delta,返回新的 count 值."""
record = self.get_for_today(user_id, usage_date, usage_type)
record.count = (record.count or 0) + delta
record.updated_at = datetime.now(timezone.utc)
self.session.flush()
return record.count
model.count += 1
model.updated_at = datetime.now(timezone.utc)
self.session.add(model)
def set_count(self, user_id: str, usage_date: date, usage_type: str, count: int) -> None:
record = self.get_for_today(user_id, usage_date, usage_type)
record.count = count
record.updated_at = datetime.now(timezone.utc)
self.session.flush()
def sync_from_redis(self, items: list[tuple[str, date, str, int]]) -> int:
"""批量将 Redis 计数同步到 DB。
items: [(user_id, usage_date, usage_type, count), ...]
返回同步的记录条数。
"""
synced = 0
for user_id, usage_date, usage_type, count in items:
self.set_count(user_id, usage_date, usage_type, count)
synced += 1
self.session.commit()
return self._to_domain(model)
@staticmethod
def _to_domain(model: DailyUsageRecordModel) -> DailyUsageRecord:
return DailyUsageRecord(
id=model.id,
user_id=model.user_id,
usage_date=model.usage_date,
usage_type=model.usage_type,
count=model.count,
updated_at=model.updated_at,
)
return synced
+45 -27
View File
@@ -1,7 +1,20 @@
from datetime import datetime, timezone
from typing import Any
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Index, Integer, String, Text, UniqueConstraint, text
from sqlalchemy import (
JSON,
Boolean,
Column,
Date,
DateTime,
Float,
Index,
Integer,
String,
Text,
UniqueConstraint,
text,
)
from sqlalchemy.orm import declarative_base
Base: Any = declarative_base()
@@ -39,11 +52,11 @@ class UserModel(Base):
phone_verified = Column(Boolean, nullable=False, default=False)
binding_completed_at = Column(DateTime, nullable=True)
profile_completed = Column(Boolean, nullable=False, default=True, server_default="true")
# 会员+积分 (#1895)
is_member = Column(Boolean, nullable=False, default=False)
member_type = Column(String(20), nullable=True)
# 会员 + 积分(#1895)
is_member = Column(Boolean, nullable=False, default=False, server_default="false")
member_type = Column(String(20), nullable=True) # monthly / quarterly / yearly
member_expires_at = Column(DateTime, nullable=True)
points_balance = Column(Integer, nullable=False, default=0)
points_balance = Column(Integer, nullable=False, default=0, server_default="0")
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -751,21 +764,25 @@ class AiAvatarRenderJob(Base):
class PointsAccountModel(Base):
"""积分账户 ORM 模型 (#1895)"""
"""积分账户(#1895) — 每用户一条记录,记录余额及累计值."""
__tablename__ = "points_accounts"
id = Column(String(36), primary_key=True)
user_id = Column(String(36), nullable=False, unique=True, index=True)
balance = Column(Integer, nullable=False, default=0)
total_earned = Column(Integer, nullable=False, default=0)
total_spent = Column(Integer, nullable=False, default=0)
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
user_id = Column(
String(36), nullable=False, unique=True, index=True
) # UNIQUE FK handled by ForeignKeyConstraint in migration
balance = Column(Integer, nullable=False, default=0, server_default="0")
total_earned = Column(Integer, nullable=False, default=0, server_default="0")
total_spent = Column(Integer, nullable=False, default=0, server_default="0")
total_purchased = Column(Integer, nullable=False, default=0, server_default="0")
total_gifted = Column(Integer, nullable=False, default=0, server_default="0")
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
class PointsTransactionModel(Base):
"""积分流水 ORM 模型 (#1895)"""
"""积分流水(#1895) — 每笔积分变动一条记录,只追加不修改."""
__tablename__ = "points_transactions"
@@ -778,38 +795,39 @@ class PointsTransactionModel(Base):
balance_after = Column(Integer, nullable=False)
description = Column(String(255), nullable=False, default="")
ref_id = Column(String(100), nullable=False, default="")
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc), index=True)
class PointsOrderModel(Base):
"""积分/会员订单 ORM 模型 (#1895)"""
"""积分充值订单(#1895)."""
__tablename__ = "points_orders"
id = Column(String(36), primary_key=True)
user_id = Column(String(36), nullable=False, index=True)
order_type = Column(String(20), nullable=False) # membership / points
product_code = Column(String(50), nullable=False)
amount_cents = Column(Integer, nullable=False)
original_amount_cents = Column(Integer, nullable=False, default=0)
package_name = Column(String(50), nullable=False)
points_amount = Column(Integer, nullable=False)
price_cents = Column(Integer, nullable=False)
currency = Column(String(10), nullable=False, default="CNY")
discount = Column(Float, nullable=False, default=1.0)
points_amount = Column(Integer, nullable=False, default=0)
original_price_cents = Column(Integer, nullable=False)
status = Column(String(20), nullable=False, default="pending", index=True)
payment_method = Column(String(50), nullable=True)
payment_id = Column(String(100), nullable=True)
paid_at = Column(DateTime, nullable=True)
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
paid_at = Column(DateTime(timezone=True), nullable=True)
expire_at = Column(DateTime(timezone=True), nullable=True)
created_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
class DailyUsageRecordModel(Base):
"""每日使用记录 ORM 模型 (#1895)"""
"""每日使用计数(#1895) — DB 持久化兜底,Redis 实时计数."""
__tablename__ = "daily_usage_records"
__table_args__ = (UniqueConstraint("user_id", "usage_date", "usage_type", name="uq_daily_usage_user_date_type"),)
id = Column(String(36), primary_key=True)
user_id = Column(String(36), nullable=False, index=True)
usage_date = Column(DateTime, nullable=False) # stored as DATE in SQL but DateTime for ORM compat
usage_type = Column(String(50), nullable=False, default="free_clip")
count = Column(Integer, nullable=False, default=0)
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
usage_date = Column(Date, nullable=False)
usage_type = Column(String(50), nullable=False)
count = Column(Integer, nullable=False, default=0, server_default="0")
updated_at = Column(DateTime(timezone=True), nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -1,55 +1,94 @@
from datetime import datetime, timezone
from __future__ import annotations
from datetime import datetime, timezone
from uuid import uuid4
from sqlalchemy import text
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import PointsAccountModel
from packages.domain.points_account import PointsAccount
class SQLAlchemyPointsAccountRepository:
"""积分账户仓储 — 纯数据库操作,无业务逻辑."""
def __init__(self, session: Session):
self.session = session
def create(self, account: PointsAccount) -> PointsAccount:
def get_by_user_id(self, user_id: str) -> PointsAccountModel | None:
return self.session.query(PointsAccountModel).filter(PointsAccountModel.user_id == user_id).first()
def create_if_not_exists(self, user_id: str) -> PointsAccountModel:
"""原子 UPSERT:按 user_id 存在则返回,否则创建余额 0 账户.
使用 with_for_update 行锁防并发重复创建;冲突(UniqueConstraint)时回退到查询。
"""
existing = self.get_by_user_id(user_id)
if existing is not None:
return existing
model = PointsAccountModel(
id=account.id,
user_id=account.user_id,
balance=account.balance,
total_earned=account.total_earned,
total_spent=account.total_spent,
created_at=account.created_at,
updated_at=account.updated_at,
id=str(uuid4()),
user_id=user_id,
balance=0,
total_earned=0,
total_spent=0,
total_purchased=0,
total_gifted=0,
created_at=datetime.now(timezone.utc),
updated_at=datetime.now(timezone.utc),
)
self.session.add(model)
self.session.commit()
return account
try:
self.session.add(model)
self.session.flush()
return model
except Exception:
self.session.rollback()
existing = self.get_by_user_id(user_id)
if existing is not None:
return existing
raise
def get_by_user_id(self, user_id: str) -> PointsAccount | None:
model = self.session.query(PointsAccountModel).filter(PointsAccountModel.user_id == user_id).first()
if model is None:
return None
return self._to_domain(model)
def get_for_update(self, user_id: str) -> PointsAccountModel | None:
"""SELECT ... FOR UPDATE,事务内锁定账户行防止并发超扣."""
return (
self.session.query(PointsAccountModel)
.filter(PointsAccountModel.user_id == user_id)
.with_for_update()
.first()
)
def update_balance(self, account: PointsAccount) -> PointsAccount:
model = self.session.query(PointsAccountModel).filter(PointsAccountModel.id == account.id).first()
def update_balance(self, account_id: str, delta: int) -> bool:
"""原子 UPDATE balance = balance + :delta.
通过 WHERE balance + :delta >= 0 保证不出现负余额;返回是否成功。
调用方负责维护 total_earned/total_spent 等累计字段(通过 update_totals)。
"""
sql = text(
"UPDATE points_accounts "
"SET balance = balance + :delta, updated_at = NOW() "
"WHERE id = :id AND balance + :delta >= 0"
)
result = self.session.execute(sql, {"delta": delta, "id": account_id})
return result.rowcount > 0
def update_totals(
self,
account_id: str,
*,
earned_delta: int = 0,
spent_delta: int = 0,
purchased_delta: int = 0,
gifted_delta: int = 0,
balance_delta: int = 0,
) -> None:
"""更新累计字段及余额(使用 Python 层 + 事务,搭配 with_for_update 使用)."""
model = self.session.get(PointsAccountModel, account_id)
if model is None:
return account
model.balance = account.balance
model.total_earned = account.total_earned
model.total_spent = account.total_spent
return
model.balance = (model.balance or 0) + balance_delta
model.total_earned = (model.total_earned or 0) + earned_delta
model.total_spent = (model.total_spent or 0) + spent_delta
model.total_purchased = (model.total_purchased or 0) + purchased_delta
model.total_gifted = (model.total_gifted or 0) + gifted_delta
model.updated_at = datetime.now(timezone.utc)
self.session.add(model)
self.session.commit()
return account
@staticmethod
def _to_domain(model: PointsAccountModel) -> PointsAccount:
return PointsAccount(
id=model.id,
user_id=model.user_id,
balance=model.balance,
total_earned=model.total_earned,
total_spent=model.total_spent,
created_at=model.created_at,
updated_at=model.updated_at,
)
self.session.flush()
@@ -1,50 +1,62 @@
from datetime import datetime
from __future__ import annotations
from datetime import datetime, timezone
from uuid import uuid4
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import PointsOrderModel
from packages.domain.points_order import PointsOrder
class SQLAlchemyPointsOrderRepository:
"""积分充值订单仓储."""
def __init__(self, session: Session):
self.session = session
def create(self, order: PointsOrder) -> PointsOrder:
def create(
self,
*,
user_id: str,
package_name: str,
points_amount: int,
price_cents: int,
original_price_cents: int,
discount: float = 1.0,
currency: str = "CNY",
payment_method: str | None = None,
expire_at: datetime | None = None,
) -> PointsOrderModel:
model = PointsOrderModel(
id=order.id,
user_id=order.user_id,
order_type=order.order_type,
product_code=order.product_code,
amount_cents=order.amount_cents,
original_amount_cents=order.original_amount_cents,
discount=order.discount,
points_amount=order.points_amount,
status=order.status,
payment_method=order.payment_method,
payment_id=order.payment_id,
paid_at=order.paid_at,
created_at=order.created_at,
id=str(uuid4()),
user_id=user_id,
package_name=package_name,
points_amount=points_amount,
price_cents=price_cents,
currency=currency,
discount=discount,
original_price_cents=original_price_cents,
status="pending",
payment_method=payment_method,
expire_at=expire_at,
created_at=datetime.now(timezone.utc),
)
self.session.add(model)
self.session.commit()
return order
self.session.flush()
return model
def get(self, order_id: str) -> PointsOrder | None:
model = self.session.query(PointsOrderModel).filter(PointsOrderModel.id == order_id).first()
if model is None:
return None
return self._to_domain(model)
def get_by_id(self, order_id: str) -> PointsOrderModel | None:
return self.session.get(PointsOrderModel, order_id)
def update_status(
self,
order_id: str,
status: str,
*,
status: str,
payment_id: str | None = None,
paid_at: datetime | None = None,
) -> PointsOrder | None:
model = self.session.query(PointsOrderModel).filter(PointsOrderModel.id == order_id).first()
) -> PointsOrderModel | None:
model = self.session.get(PointsOrderModel, order_id)
if model is None:
return None
model.status = status
@@ -52,45 +64,19 @@ class SQLAlchemyPointsOrderRepository:
model.payment_id = payment_id
if paid_at is not None:
model.paid_at = paid_at
self.session.add(model)
self.session.commit()
return self._to_domain(model)
self.session.flush()
return model
def list_by_user(
self,
user_id: str,
*,
order_type: str | None = None,
status: str | None = None,
page: int = 1,
page_size: int = 20,
) -> tuple[list[PointsOrder], int]:
query = self.session.query(PointsOrderModel).filter(PointsOrderModel.user_id == user_id)
if order_type:
query = query.filter(PointsOrderModel.order_type == order_type)
if status:
query = query.filter(PointsOrderModel.status == status)
total = query.count()
models = (
query.order_by(PointsOrderModel.created_at.desc()).offset((page - 1) * page_size).limit(page_size).all()
)
return [self._to_domain(m) for m in models], total
@staticmethod
def _to_domain(model: PointsOrderModel) -> PointsOrder:
return PointsOrder(
id=model.id,
user_id=model.user_id,
order_type=model.order_type,
product_code=model.product_code,
amount_cents=model.amount_cents,
original_amount_cents=model.original_amount_cents,
discount=model.discount,
points_amount=model.points_amount,
status=model.status,
payment_method=model.payment_method,
payment_id=model.payment_id,
paid_at=model.paid_at,
created_at=model.created_at,
def list_expired_pending(self, before: datetime | None = None) -> list[PointsOrderModel]:
"""查出 expire_at 已过仍处于 pending 状态的订单(可用于定时取消)."""
if before is None:
before = datetime.now(timezone.utc)
return (
self.session.query(PointsOrderModel)
.filter(
PointsOrderModel.status == "pending",
PointsOrderModel.expire_at.isnot(None),
PointsOrderModel.expire_at < before,
)
.all()
)
@@ -1,65 +1,84 @@
from __future__ import annotations
from datetime import datetime, timezone
from typing import Optional
from uuid import uuid4
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import PointsTransactionModel
from packages.domain.points_transaction import PointsTransaction
class SQLAlchemyPointsTransactionRepository:
"""积分流水仓储 — 只追加,不修改/删除."""
def __init__(self, session: Session):
self.session = session
def create(self, transaction: PointsTransaction) -> PointsTransaction:
def create(
self,
*,
user_id: str,
account_id: str,
type_: str,
source: str,
amount: int,
balance_after: int,
description: str = "",
ref_id: str = "",
) -> PointsTransactionModel:
model = PointsTransactionModel(
id=transaction.id,
user_id=transaction.user_id,
account_id=transaction.account_id,
type=transaction.type,
source=transaction.source,
amount=transaction.amount,
balance_after=transaction.balance_after,
description=transaction.description,
ref_id=transaction.ref_id,
created_at=transaction.created_at,
id=str(uuid4()),
user_id=user_id,
account_id=account_id,
type=type_,
source=source,
amount=amount,
balance_after=balance_after,
description=description,
ref_id=ref_id,
created_at=datetime.now(timezone.utc),
)
self.session.add(model)
self.session.commit()
return transaction
self.session.flush()
return model
def get_by_id(self, tx_id: str) -> PointsTransactionModel | None:
return self.session.get(PointsTransactionModel, tx_id)
def exists_refund_for(self, original_tx_id: str) -> bool:
"""判断给定原 spend 流水是否已有 refund 流水(幂等检查)."""
from sqlalchemy import func
return bool(
self.session.query(func.count(PointsTransactionModel.id))
.filter(
PointsTransactionModel.type == "refund",
PointsTransactionModel.ref_id == original_tx_id,
)
.scalar()
)
def list_by_user(
self,
user_id: str,
*,
type: str | None = None,
source: str | None = None,
page: int = 1,
page_size: int = 20,
) -> tuple[list[PointsTransaction], int]:
query = self.session.query(PointsTransactionModel).filter(PointsTransactionModel.user_id == user_id)
if type:
query = query.filter(PointsTransactionModel.type == type)
offset: int = 0,
limit: int = 20,
type_: Optional[str] = None,
source: Optional[str] = None,
start_date: Optional[datetime] = None,
end_date: Optional[datetime] = None,
) -> tuple[list[PointsTransactionModel], int]:
q = self.session.query(PointsTransactionModel).filter(PointsTransactionModel.user_id == user_id)
if type_:
q = q.filter(PointsTransactionModel.type == type_)
if source:
query = query.filter(PointsTransactionModel.source == source)
total = query.count()
models = (
query.order_by(PointsTransactionModel.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
.all()
)
return [self._to_domain(m) for m in models], total
@staticmethod
def _to_domain(model: PointsTransactionModel) -> PointsTransaction:
return PointsTransaction(
id=model.id,
user_id=model.user_id,
account_id=model.account_id,
type=model.type,
source=model.source,
amount=model.amount,
balance_after=model.balance_after,
description=model.description or "",
ref_id=model.ref_id or "",
created_at=model.created_at,
)
q = q.filter(PointsTransactionModel.source == source)
if start_date:
q = q.filter(PointsTransactionModel.created_at >= start_date)
if end_date:
q = q.filter(PointsTransactionModel.created_at <= end_date)
total = q.count()
items = q.order_by(PointsTransactionModel.created_at.desc()).offset(offset).limit(limit).all()
return items, total
@@ -30,6 +30,10 @@ class SQLAlchemyUserRepository(UserRepository):
model.subscription_plan = user.subscription_plan
model.subscription_status = user.subscription_status
model.subscription_expires_at = user.subscription_expires_at
model.is_member = user.is_member
model.member_type = user.member_type
model.member_expires_at = user.member_expires_at
model.points_balance = user.points_balance
model.max_projects = user.max_projects
model.max_storage_gb = user.max_storage_gb
model.is_admin = user.is_admin
@@ -106,6 +110,10 @@ class SQLAlchemyUserRepository(UserRepository):
subscription_plan=model.subscription_plan or "free",
subscription_status=model.subscription_status or "active",
subscription_expires_at=model.subscription_expires_at,
is_member=bool(model.is_member or False),
member_type=model.member_type,
member_expires_at=model.member_expires_at,
points_balance=int(model.points_balance or 0),
max_projects=model.max_projects or 3,
max_storage_gb=model.max_storage_gb or 10,
is_admin=model.is_admin or False,
+604
View File
@@ -0,0 +1,604 @@
"""PointsService — 会员积分应用服务(#1895 P1+P2).
核心能力:账户查询、积分扣减/退还/充值、充值订单管理、每日免费额度(Redis 计数 + DB 兜底)。
所有 DB 操作通过 repository 层;事务通过 session 上下文保证原子性。
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from datetime import date, datetime, timedelta, timezone
from typing import Any
from packages.adapters.sqlalchemy_impl.daily_usage_repository import (
SQLAlchemyDailyUsageRepository,
)
from packages.adapters.sqlalchemy_impl.points_account_repository import (
SQLAlchemyPointsAccountRepository,
)
from packages.adapters.sqlalchemy_impl.points_order_repository import (
SQLAlchemyPointsOrderRepository,
)
from packages.adapters.sqlalchemy_impl.points_transaction_repository import (
SQLAlchemyPointsTransactionRepository,
)
from packages.domain.points import (
FREE_DAILY_CLIPS,
ORDER_STATUS_PAID,
ORDER_STATUS_PENDING,
POINTS_PACKAGES,
TX_SOURCE_RECHARGE,
TX_TYPE_EARN,
TX_TYPE_REFUND,
TX_TYPE_SPEND,
calc_package_price,
calc_points,
)
logger = logging.getLogger(__name__)
CST = timezone(timedelta(hours=8))
DAILY_USAGE_REDIS_PREFIX = "daily_usage"
DAILY_USAGE_REDIS_TTL_SECONDS = 48 * 3600 # 48h
# 会员套餐与权益
MEMBERSHIP_PLANS: dict[str, dict[str, Any]] = {
"monthly": {
"name": "月度会员",
"days": 31,
"price_cents": 2900,
"discount": 0.9,
"daily_free_clips_limit": -1, # -1 表示不限
},
"quarterly": {
"name": "季度会员",
"days": 93,
"price_cents": 7900,
"discount": 0.87,
"daily_free_clips_limit": -1,
},
"yearly": {
"name": "年度会员",
"days": 366,
"price_cents": 25900,
"discount": 0.8,
"daily_free_clips_limit": -1,
},
}
MEMBER_PACKAGE_NAME_PREFIX = "membership_" # 订单 package_name 前缀,用于区分会员订阅单
POINTS_UNIT_PRICE_YUAN = 0.10 # 积分单价(元/积分),用于展示
@dataclass(slots=True)
class DeductResult:
success: bool
balance: int
amount: int = 0
transaction_id: str | None = None
balance_after: int = 0
reason: str = ""
@dataclass(slots=True)
class EarnResult:
success: bool
transaction_id: str | None = None
balance_after: int = 0
amount: int = 0
reason: str = ""
@dataclass(slots=True)
class DailyUsageInfo:
used: int
limit: int
remaining: int
reset_at: datetime | None = None
class PointsService:
"""积分业务服务。
依赖通过构造函数注入(repositories + 可选 redis 客户端)。
事务:使用 repo.session.begin() 上下文保证原子性,失败回滚。
"""
def __init__(
self,
account_repo: SQLAlchemyPointsAccountRepository,
tx_repo: SQLAlchemyPointsTransactionRepository,
order_repo: SQLAlchemyPointsOrderRepository,
daily_usage_repo: SQLAlchemyDailyUsageRepository,
redis_client: Any | None = None,
):
self._accounts = account_repo
self._txs = tx_repo
self._orders = order_repo
self._daily = daily_usage_repo
self._redis = redis_client
# ── 账户 ──────────────────────────────────────────────────────────────
def get_account(self, user_id: str):
"""获取积分账户,不存在则自动创建(余额 0)."""
return self._accounts.create_if_not_exists(user_id)
def get_balance(self, user_id: str) -> int:
account = self._accounts.get_by_user_id(user_id)
if account is None:
account = self._accounts.create_if_not_exists(user_id)
self._accounts.session.commit()
return int(account.balance or 0)
# ── 纯计算 ────────────────────────────────────────────────────────────
def calculate_cost(
self,
scene_key: str,
is_member: bool,
duration_minutes: float = 1.0,
extra_segments: int = 0,
) -> int:
"""仅计算所需积分(不扣减、不写库),用于前端预估消耗。"""
return calc_points(
scene_key,
is_member,
duration_minutes=duration_minutes,
extra_segments=extra_segments,
)
# ── 流水查询 ──────────────────────────────────────────────────────────
def list_transactions(
self,
user_id: str,
*,
offset: int = 0,
limit: int = 20,
type_: str | None = None,
source: str | None = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
):
"""分页查询积分流水,返回 (items, total)。"""
return self._txs.list_by_user(
user_id,
offset=offset,
limit=limit,
type_=type_,
source=source,
start_date=start_date,
end_date=end_date,
)
# ── 会员订阅 ──────────────────────────────────────────────────────────
def create_membership_order(
self,
user_id: str,
member_type: str,
payment_method: str,
):
"""创建会员订阅订单(pending 状态,复用 points_orders 表,package_name 前缀区分)。"""
plan = MEMBERSHIP_PLANS.get(member_type)
if plan is None:
raise ValueError(f"未知会员类型: {member_type}")
expire_at = datetime.now(timezone.utc) + timedelta(minutes=30)
order = self._orders.create(
user_id=user_id,
package_name=f"{MEMBER_PACKAGE_NAME_PREFIX}{member_type}",
points_amount=0,
price_cents=int(plan["price_cents"]),
original_price_cents=int(plan["price_cents"]),
discount=float(plan["discount"]),
currency="CNY",
payment_method=payment_method,
expire_at=expire_at,
)
self._orders.session.commit()
return order
def subscribe_member(self, user_id: str, member_type: str) -> None:
"""激活/续费会员:更新 users.is_member/member_type/member_expires_at。
新会员从当前时间加 days;已在会员期内则在 member_expires_at 基础上顺延。
调用方需确保在事务中调用。
"""
from packages.adapters.sqlalchemy_impl.models import UserModel
plan = MEMBERSHIP_PLANS.get(member_type)
if plan is None:
raise ValueError(f"未知会员类型: {member_type}")
session = self._orders.session
user_model = session.query(UserModel).filter(UserModel.id == user_id).with_for_update().first()
if user_model is None:
raise ValueError(f"用户不存在: {user_id}")
now = datetime.now(timezone.utc)
base = (
user_model.member_expires_at
if (user_model.member_expires_at and user_model.member_expires_at > now)
else now
)
new_expires = base + timedelta(days=int(plan["days"]))
user_model.is_member = True
user_model.member_type = member_type
user_model.member_expires_at = new_expires
session.flush()
# ── 扣减 / 退还 ───────────────────────────────────────────────────────
def check_and_deduct(
self,
user_id: str,
scene_key: str,
duration_minutes: float = 1.0,
extra_segments: int = 0,
ref_id: str = "",
description: str = "",
is_member: bool = False,
) -> DeductResult:
"""事务性扣减积分。
1. calc_points 计算消耗量
2. SELECT FOR UPDATE 锁定账户
3. 余额不足返回失败(reason=insufficient)
4. 余额充足:更新 balance/total_spent,插入 spend 流水,提交事务
"""
amount = calc_points(
scene_key,
is_member,
duration_minutes=duration_minutes,
extra_segments=extra_segments,
)
if amount <= 0:
# 免费场景(如 voice_clone_train)直接返回成功
account = self._accounts.create_if_not_exists(user_id)
self._accounts.session.commit()
return DeductResult(
success=True,
balance=int(account.balance or 0),
amount=0,
transaction_id=None,
balance_after=int(account.balance or 0),
reason="free_scene",
)
session = self._accounts.session
with session.begin_nested() if session.in_transaction() else session.begin():
account = self._accounts.get_for_update(user_id)
if account is None:
account = self._accounts.create_if_not_exists(user_id)
session.flush()
current = int(account.balance or 0)
if current < amount:
return DeductResult(
success=False,
balance=current,
amount=amount,
reason="insufficient",
)
account.balance = current - amount
account.total_spent = int(account.total_spent or 0) + amount
account.updated_at = datetime.now(timezone.utc)
session.flush()
tx = self._txs.create(
user_id=user_id,
account_id=account.id,
type_=TX_TYPE_SPEND,
source=scene_key,
amount=amount,
balance_after=int(account.balance or 0),
description=description,
ref_id=ref_id,
)
return DeductResult(
success=True,
balance=int(account.balance or 0),
amount=amount,
transaction_id=tx.id,
balance_after=int(account.balance or 0),
)
def refund(self, user_id: str, transaction_id: str, reason: str = "") -> bool:
"""根据原 spend 流水退还积分。
- 只能退还 type=spend 的流水(防止重复退还 earn)
- 事务内 balance+amount, total_spent-amount,插入 refund 流水
"""
session = self._accounts.session
with session.begin_nested() if session.in_transaction() else session.begin():
orig = self._txs.get_by_id(transaction_id)
if orig is None:
logger.warning("refund: 流水不存在 tx_id=%s", transaction_id)
return False
if orig.type != TX_TYPE_SPEND:
logger.warning(
"refund: 非 spend 流水不可退 tx_id=%s type=%s",
transaction_id,
orig.type,
)
return False
if self._txs.exists_refund_for(transaction_id):
logger.warning("refund: 流水已退款 tx_id=%s", transaction_id)
return False
account = self._accounts.get_for_update(user_id)
if account is None or account.id != orig.account_id:
logger.warning(
"refund: 账户不匹配 user_id=%s tx_account=%s",
user_id,
orig.account_id,
)
return False
amount = int(orig.amount or 0)
account.balance = int(account.balance or 0) + amount
account.total_spent = max(0, int(account.total_spent or 0) - amount)
account.updated_at = datetime.now(timezone.utc)
session.flush()
self._txs.create(
user_id=user_id,
account_id=account.id,
type_=TX_TYPE_REFUND,
source=orig.source,
amount=amount,
balance_after=int(account.balance or 0),
description=reason or f"refund for {transaction_id}",
ref_id=transaction_id,
)
return True
# ── 充值 / 奖励 ───────────────────────────────────────────────────────
def earn_points(
self,
user_id: str,
amount: int,
source: str,
description: str = "",
ref_id: str = "",
) -> EarnResult:
"""获得积分(充值或任务奖励)。
source=recharge 累加 total_purchased;source=task_reward 累加 total_gifted。
"""
if amount <= 0:
return EarnResult(success=False, reason="invalid_amount")
session = self._accounts.session
with session.begin_nested() if session.in_transaction() else session.begin():
account = self._accounts.get_for_update(user_id)
if account is None:
account = self._accounts.create_if_not_exists(user_id)
session.flush()
account.balance = int(account.balance or 0) + amount
account.total_earned = int(account.total_earned or 0) + amount
if source == TX_SOURCE_RECHARGE:
account.total_purchased = int(account.total_purchased or 0) + amount
else:
account.total_gifted = int(account.total_gifted or 0) + amount
account.updated_at = datetime.now(timezone.utc)
session.flush()
tx = self._txs.create(
user_id=user_id,
account_id=account.id,
type_=TX_TYPE_EARN,
source=source,
amount=amount,
balance_after=int(account.balance or 0),
description=description,
ref_id=ref_id,
)
return EarnResult(
success=True,
transaction_id=tx.id,
balance_after=int(account.balance or 0),
amount=amount,
)
# ── 订单 ──────────────────────────────────────────────────────────────
def create_order(
self,
user_id: str,
package_id: str,
payment_method: str,
member_type_for_discount: str | None = None,
):
"""创建积分充值订单(pending 状态,30 分钟过期)。"""
discounted_cents, original_cents, discount = calc_package_price(package_id, member_type_for_discount)
pkg = next((p for p in POINTS_PACKAGES if p["id"] == package_id), None)
if pkg is None:
raise ValueError(f"未知积分包: {package_id}")
expire_at = datetime.now(timezone.utc) + timedelta(minutes=30)
order = self._orders.create(
user_id=user_id,
package_name=pkg["name"],
points_amount=int(pkg["points"]),
price_cents=int(discounted_cents),
original_price_cents=int(original_cents),
discount=float(discount),
currency="CNY",
payment_method=payment_method,
expire_at=expire_at,
)
self._orders.session.commit()
return order
def mark_order_paid(self, order_id: str, payment_id: str):
"""标记订单已支付:积分包 → 充值积分;会员订阅单 → 激活/续费会员。
事务内:改订单状态为 paid -> 分发给对应的履约逻辑。幂等。
"""
session = self._orders.session
with session.begin_nested() if session.in_transaction() else session.begin():
order = self._orders.get_by_id(order_id)
if order is None:
raise ValueError(f"订单不存在: {order_id}")
if order.status == ORDER_STATUS_PAID:
return order # 幂等
if order.status != ORDER_STATUS_PENDING:
raise ValueError(f"订单状态不可支付: {order.status}")
paid_at = datetime.now(timezone.utc)
self._orders.update_status(
order_id,
status=ORDER_STATUS_PAID,
payment_id=payment_id,
paid_at=paid_at,
)
if order.package_name and order.package_name.startswith(MEMBER_PACKAGE_NAME_PREFIX):
member_type = order.package_name[len(MEMBER_PACKAGE_NAME_PREFIX) :]
self.subscribe_member(order.user_id, member_type)
else:
# earn_points 在同一事务内(使用 orders 的 session 需要重新获取账户 repo
# 为了保证在同一事务,我们让 earn_points 通过 begin_nested 使用;
# 注意: account/tx repo 与 order repo 应共享同一 session
# 此处调用 earn_points 将在 order 事务内做 SAVEPOINT
self.earn_points(
user_id=order.user_id,
amount=int(order.points_amount),
source=TX_SOURCE_RECHARGE,
description=f"充值 {order.package_name}",
ref_id=order_id,
)
session.flush()
# commit 外层事务
session.commit()
return self._orders.get_by_id(order_id)
# ── 每日免费额度(Redis 主 + DB 兜底) ─────────────────────────────────
@staticmethod
def _today_cst() -> date:
return datetime.now(CST).date()
@classmethod
def _redis_key(cls, user_id: str, usage_date: date, usage_type: str) -> str:
return f"{DAILY_USAGE_REDIS_PREFIX}:{user_id}:{usage_date.strftime('%Y%m%d')}:{usage_type}"
def _redis_check_and_incr(self, key: str, limit: int) -> tuple[bool, int] | None:
"""原子 check+incr:若当前值已 >= limit 则不递增(拒绝);否则 +1。
使用 Lua 脚本保证原子性,避免超限时仍被计数导致"被占用"额度。
返回 (allowed: bool, current_count_after: int);Redis 不可用时返回 None。
- allowed=True 表示本次占用成功,count 为占用后的次数(1..limit)
- allowed=False 表示已达上限,count 为已占用次数(=limit)
"""
if self._redis is None:
return None
try:
# 返回: {0=rejected, 1=allowed}, current_count
lua = """
local cur = tonumber(redis.call('GET', KEYS[1]) or '0')
local lim = tonumber(ARGV[1])
if cur >= lim then
return {0, tostring(cur)}
end
local nv = redis.call('INCR', KEYS[1])
if tonumber(nv) == 1 then
redis.call('EXPIRE', KEYS[1], tonumber(ARGV[2]))
end
return {1, tostring(nv)}
"""
allowed_flag, cur_str = self._redis.eval(lua, 1, key, limit, DAILY_USAGE_REDIS_TTL_SECONDS)
return bool(int(allowed_flag)), int(cur_str)
except Exception:
logger.warning("Redis check+incr 失败 key=%s", key, exc_info=True)
return None
def _redis_get(self, key: str) -> int | None:
if self._redis is None:
return None
try:
v = self._redis.get(key)
return int(v) if v is not None else 0
except Exception:
logger.warning("Redis GET 失败 key=%s", key, exc_info=True)
return None
def check_and_incr_daily_free_clips(
self,
user_id: str,
usage_type: str = "free_clip",
limit: int = FREE_DAILY_CLIPS,
) -> bool:
"""检查并占用一次每日免费额度。
- 优先走 Redis INCR(原子 + TTL 48h)
- Redis 不可用则降级 DB 行锁 + increment_count
- 超过 limit 返回 False;否则 True
- 异步将 Redis 计数同步到 DB(同步写入,简单可靠;后续可改为异步任务)
"""
today = self._today_cst()
key = self._redis_key(user_id, today, usage_type)
result = self._redis_check_and_incr(key, limit)
if result is not None:
allowed, current = result
# 异步同步到 DB(这里同步写,轻量)
try:
self._daily.set_count(user_id, today, usage_type, current)
self._daily.session.commit()
except Exception:
logger.warning("daily_usage DB 同步失败 user=%s", user_id, exc_info=True)
self._daily.session.rollback()
return allowed
# Redis 不可用:走 DB(事务内先查后增,防超限)
session = self._daily.session
try:
with session.begin_nested() if session.in_transaction() else session.begin():
record = self._daily.get_for_today(user_id, today, usage_type)
current = int(record.count or 0)
if current >= limit:
session.rollback()
# 即便回滚本次嵌套事务,仍保留会话可用,提交前序
return False
self._daily.increment_count(user_id, today, usage_type)
session.commit()
return True
except Exception:
session.rollback()
raise
def get_daily_usage(
self,
user_id: str,
usage_type: str = "free_clip",
limit: int = FREE_DAILY_CLIPS,
) -> DailyUsageInfo:
"""查询今日免费额度使用情况."""
today = self._today_cst()
key = self._redis_key(user_id, today, usage_type)
used = self._redis_get(key)
if used is None:
# 从 DB 读
record = self._daily.get_for_today(user_id, today, usage_type)
try:
self._daily.session.commit()
except Exception:
self._daily.session.rollback()
used = int(record.count or 0)
remaining = max(0, limit - used)
# reset_at: 次日 00:00 CST
tomorrow_cst = datetime.combine(today + timedelta(days=1), datetime.min.time(), tzinfo=CST)
return DailyUsageInfo(
used=used,
limit=limit,
remaining=remaining,
reset_at=tomorrow_cst,
)
def list_points_packages(member_type_for_discount: str | None = None) -> list[dict[str, Any]]:
"""返回积分包列表(含按会员类型计算的折后价)。无状态工具函数。"""
items: list[dict[str, Any]] = []
for pkg in POINTS_PACKAGES:
discounted_cents, original_cents, discount = calc_package_price(pkg["id"], member_type_for_discount)
items.append(
{
"id": pkg["id"],
"name": pkg["name"],
"points": pkg["points"],
"price_cents": original_cents,
"discounted_price_cents": discounted_cents,
"discount": discount,
}
)
return items
+8
View File
@@ -110,6 +110,10 @@ class APISettings(SharedSettings):
# 渲染引擎选择:legacy=旧VideoComposeService,unified=新UnifiedRenderService
render_engine: str = "legacy"
# ── 会员积分系统开关 ────────────────────────────────────────────────
# 默认关闭,开发/测试期不实际扣费;上线时通过 POINTS_ENABLED=true 开启
points_enabled: bool = False
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
@@ -280,6 +284,10 @@ class APISettings(SharedSettings):
def RENDER_ENGINE(self) -> str:
return self.render_engine
@property
def POINTS_ENABLED(self) -> bool:
return self.points_enabled
def get_api_settings() -> APISettings:
"""获取 API 配置单例(统一入口)。"""
+6 -1
View File
@@ -39,9 +39,14 @@ class User:
last_login_at: datetime | None = None
last_login_ip: str | None = None
# 订阅相关字段 (移到 User 级别)
subscription_plan: str = "free" # free, pro, enterprise
subscription_plan: str = "free" # free, pro, enterprise(旧字段,新逻辑使用 is_member/member_type)
subscription_status: str = "active" # active, cancelled, expired
subscription_expires_at: datetime | None = None
# 会员&积分字段(#1895 P1)
is_member: bool = False
member_type: str | None = None # monthly / quarterly / yearly
member_expires_at: datetime | None = None
points_balance: int = 0
# 配额限制 (移到 User 级别)
max_projects: int = 3 # free: 3, pro: unlimited, enterprise: unlimited
max_storage_gb: int = 10 # free: 10, pro: 100, enterprise: 1000
+254
View File
@@ -0,0 +1,254 @@
"""会员积分领域层(#1895) — 纯函数、实体、规则常量."""
from __future__ import annotations
import math
from dataclasses import dataclass
from datetime import datetime
from typing import Any
# ── 规则常量 ──────────────────────────────────────────────────────────────────
# 积分消耗场景:scene_key -> {name, base_points, unit, extra_per_30s}
# extra_per_30s: 仅 ai_video 使用——每额外 30s 多扣的积分(基础 3 积分 = ≤30s)
POINTS_RULES: dict[str, dict[str, Any]] = {
"ai_voice": {
"name": "AI 配音",
"base_points": 1,
"unit": "分钟",
"extra_per_30s": 0,
},
"ai_video": {
"name": "智能混剪",
"base_points": 3,
"unit": "条",
"extra_per_30s": 1, # 每超出 30s 多 1 积分
},
"ai_digital_human": {
"name": "AI 数字人",
"base_points": 15,
"unit": "分钟",
"extra_per_30s": 0,
},
"voice_clone_train": {
"name": "声音克隆训练",
"base_points": 0,
"unit": "次",
"extra_per_30s": 0,
},
"voice_clone_synth": {
"name": "声音克隆合成",
"base_points": 1,
"unit": "分钟",
"extra_per_30s": 0,
},
"douyin_extract": {
"name": "抖音链接提取",
"base_points": 1,
"unit": "次",
"extra_per_30s": 0,
},
"ai_rewrite": {
"name": "AI 改写文案",
"base_points": 1,
"unit": "次",
"extra_per_30s": 0,
},
"ai_title": {
"name": "AI 标题生成",
"base_points": 1,
"unit": "次",
"extra_per_30s": 0,
},
"ai_cover": {
"name": "AI 封面生成",
"base_points": 1,
"unit": "张",
"extra_per_30s": 0,
},
}
# 免费用户积分消耗倍率
FREE_USER_MULTIPLIER: float = 1.15
# 免费用户每日免费混剪条数
FREE_DAILY_CLIPS: int = 2
# 积分包(id/名称/积分数量/原价(分))
POINTS_PACKAGES: list[dict[str, Any]] = [
{"id": "starter_pack", "name": "体验包", "points": 100, "price_cents": 990},
{"id": "basic_pack", "name": "基础包", "points": 500, "price_cents": 3900},
{"id": "pro_pack", "name": "专业包", "points": 2000, "price_cents": 12900},
]
# 付费会员类型对应的积分包折扣(月/季/年)
MEMBER_PACKAGE_DISCOUNT: dict[str, float] = {
"monthly": 0.9,
"quarterly": 0.87,
"yearly": 0.8,
}
# 流水类型
TX_TYPE_EARN = "earn"
TX_TYPE_SPEND = "spend"
TX_TYPE_REFUND = "refund"
# 流水来源
TX_SOURCE_RECHARGE = "recharge"
TX_SOURCE_TASK_REWARD = "task_reward"
# 订单状态
ORDER_STATUS_PENDING = "pending"
ORDER_STATUS_PAID = "paid"
ORDER_STATUS_FAILED = "failed"
ORDER_STATUS_REFUNDED = "refunded"
# ── 纯函数 ────────────────────────────────────────────────────────────────────
def calc_points(
scene_key: str,
is_member: bool,
duration_minutes: float = 1.0,
extra_segments: int = 0,
) -> int:
"""根据场景、会员身份、时长计算本次消耗的积分(整数)。
- 会员:按 base_points + extra 计算
- 免费用户:会员价 × 1.15,向上取整
- voice_clone_train 免费,返回 0
"""
rule = POINTS_RULES.get(scene_key)
if rule is None:
raise ValueError(f"未知积分场景: {scene_key}")
base = rule["base_points"]
if base == 0:
return 0
extra_per_30s = rule.get("extra_per_30s", 0)
# ai_video 场景:duration_minutes 视为分钟数,按"每 30s"加 extra
# 基础 3 分对应 ≤30s;每多 30s 加 1 分
if scene_key == "ai_video" and extra_per_30s > 0:
# duration_minutes<=0.5 视为 0 段额外;否则每 30s 一段
extra_count = max(0, math.ceil(duration_minutes * 2) - 1)
member_points = base + extra_per_30s * extra_count
else:
# 按分钟计费的场景(ai_voice/ai_digital_human/voice_clone_synth)向上取整到分钟
if rule.get("unit") == "分钟":
minutes = max(1, math.ceil(duration_minutes))
member_points = base * minutes
else:
member_points = base
# 兼容额外段(预留)
if extra_segments > 0 and extra_per_30s > 0:
member_points += extra_per_30s * extra_segments
if is_member:
return max(0, int(member_points))
return max(0, math.ceil(member_points * FREE_USER_MULTIPLIER))
def get_package(package_id: str) -> dict[str, Any]:
for pkg in POINTS_PACKAGES:
if pkg["id"] == package_id:
return pkg
raise ValueError(f"未知积分包: {package_id}")
def calc_package_price(package_id: str, member_type_for_discount: str | None = None) -> tuple[int, int, float]:
"""计算积分包实际应付价格。
返回 (discounted_price_cents, original_price_cents, discount)。
"""
pkg = get_package(package_id)
original = int(pkg["price_cents"])
discount = 1.0
if member_type_for_discount:
discount = MEMBER_PACKAGE_DISCOUNT.get(member_type_for_discount, 1.0)
discounted = int(round(original * discount))
return discounted, original, discount
# ── 实体 ──────────────────────────────────────────────────────────────────────
@dataclass(slots=True)
class PointsAccount:
id: str
user_id: str
balance: int = 0
total_earned: int = 0
total_spent: int = 0
total_purchased: int = 0
total_gifted: int = 0
created_at: datetime | None = None
updated_at: datetime | None = None
@dataclass(slots=True)
class PointsTransaction:
id: str
user_id: str
account_id: str
type: str # earn / spend / refund
source: str
amount: int
balance_after: int
description: str = ""
ref_id: str = ""
created_at: datetime | None = None
@dataclass(slots=True)
class PointsOrder:
id: str
user_id: str
package_name: str
points_amount: int
price_cents: int
original_price_cents: int
currency: str = "CNY"
discount: float = 1.0
status: str = "pending"
payment_method: str | None = None
payment_id: str | None = None
paid_at: datetime | None = None
expire_at: datetime | None = None
created_at: datetime | None = None
@dataclass(slots=True)
class DailyUsageRecord:
id: str
user_id: str
usage_date: Any # date
usage_type: str
count: int = 0
updated_at: datetime | None = None
@dataclass(slots=True)
class DeductResult:
success: bool
balance: int
amount: int = 0
transaction_id: str | None = None
balance_after: int = 0
reason: str = "" # "insufficient" 等失败原因
@dataclass(slots=True)
class EarnResult:
success: bool
transaction_id: str | None = None
balance_after: int = 0
amount: int = 0
reason: str = ""
@dataclass(slots=True)
class DailyUsageInfo:
used: int
limit: int
remaining: int
reset_at: datetime | None = None
+136 -153
View File
@@ -1,181 +1,164 @@
"""AI 功能入口的积分扣费装饰器 (#1895)
"""PointsGate — AI 路由积分扣费中间件(#1895 P4)。
支持 sync 和 async 函数。业务失败时自动退还积分。
使用方式::
from packages.middleware.points_gate import points_deduction
@router.post("/generate")
def generate_title(
request: TitleRequest,
current_user: User = Depends(get_current_user),
points_svc: PointsService = Depends(get_points_service),
):
with points_deduction(points_svc, current_user, "ai_title", description="AI 标题生成"):
result = do_generate(...)
return result
- 默认受 ``settings.POINTS_ENABLED`` 开关控制,关闭时不扣费,直接放行。
- ``ai_video`` 场景会优先走每日免费混剪额度(免费用户每日 2 条),会员直接走扣费。
- contextmanager 内业务异常会自动 refund;未抛异常视为成功,积分正常扣除。
"""
from __future__ import annotations
import asyncio
import functools
import inspect
import logging
from typing import Any, Callable
from contextlib import contextmanager
from datetime import datetime, timezone
from typing import Any, Iterator
from fastapi import HTTPException
from fastapi import HTTPException, status
logger = logging.getLogger(__name__)
def points_gate(
def _is_active_member(user: Any) -> bool:
"""判断用户是否为「在有效期内」的付费会员."""
if not user:
return False
if not bool(getattr(user, "is_member", False)):
return False
expires_at = getattr(user, "member_expires_at", None)
if expires_at is None:
return True
now = datetime.now(timezone.utc)
if getattr(expires_at, "tzinfo", None) is None:
now = now.replace(tzinfo=None)
return expires_at > now
def _insufficient_points(amount: int, balance: int, scene_name: str = "") -> HTTPException:
return HTTPException(
status_code=status.HTTP_402_PAYMENT_REQUIRED,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {amount} 积分,当前余额 {balance}",
"required_points": amount,
"current_balance": balance,
"scene": scene_name,
},
)
@contextmanager
def points_deduction(
points_svc: Any,
user: Any,
scene_key: str,
per_unit: int | None = None,
unit_field: str | None = None,
quantity_field: str | None = None,
) -> Callable:
"""AI 功能入口积分扣费装饰器。
*,
duration_minutes: float = 1.0,
extra_segments: int = 0,
description: str = "",
enabled: bool | None = None,
) -> Iterator[str | None]:
"""积分扣费 contextmanager:业务成功 → 确认扣费;业务抛异常 → 自动退款.
Args:
scene_key: 消耗场景标识(对应 points_rules.POINTS_SCENES 的 key)
per_unit: 固定消耗积分(直接指定,不走规则计算)
unit_field: 从 request body 取时长字段名(按时长计费场景)
quantity_field: 从 request body 取数量字段名(按次计费场景)
使用示例::
@router.post("/ai/voice")
@points_gate("ai_voice", unit_field="duration_minutes")
async def create_ai_voice(body: VoiceRequest, current_user=Depends(get_current_user), db=Depends(get_db_session)):
...
:param points_svc: PointsService 实例
:param user: User entity(需带 id/is_member/member_expires_at)
:param scene_key: POINTS_RULES 中的场景 key
:param duration_minutes: 时长(分钟),按场景语义解释
:param extra_segments: 额外片段数(ai_video 预留)
:param description: 流水描述
:param enabled: 显式开关;None 时读取 settings.POINTS_ENABLED
:yields: transaction_id 或 None(免费/未启用场景)
"""
if enabled is None:
from app.config import settings
def decorator(func: Callable) -> Callable:
is_async = asyncio.iscoroutinefunction(func)
enabled = bool(getattr(settings, "POINTS_ENABLED", False))
if not enabled:
yield None
return
@functools.wraps(func)
async def async_wrapper(*args: Any, **kwargs: Any) -> Any:
return await _execute_with_gate(
func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, is_async=True
)
from packages.domain.points import POINTS_RULES, calc_points
@functools.wraps(func)
def sync_wrapper(*args: Any, **kwargs: Any) -> Any:
return _execute_with_gate(
func, args, kwargs, scene_key, per_unit, unit_field, quantity_field, is_async=False
)
if scene_key not in POINTS_RULES:
logger.debug("points_gate: 未配置场景 scene=%s,放行", scene_key)
yield None
return
if is_async:
return async_wrapper
return sync_wrapper
user_id = getattr(user, "id", None)
if not user_id:
yield None
return
return decorator
is_member = _is_active_member(user)
scene_name = POINTS_RULES[scene_key].get("name", scene_key)
desc = description or scene_name
if scene_key == "ai_video" and not is_member:
try:
if points_svc.check_and_incr_daily_free_clips(user_id):
logger.info("points_gate: 免费混剪额度占用 user=%s", user_id)
yield None
return
except Exception:
logger.warning("points_gate: daily_free_clips 检查失败,降级走扣费 user=%s", user_id, exc_info=True)
def _extract_kwargs(func: Callable, args: tuple, kwargs: dict) -> dict:
"""将位置参数映射到函数签名中的参数名,便于统一按 kwargs 提取。"""
sig = inspect.signature(func)
bound = sig.bind_partial(*args, **kwargs)
merged = dict(bound.arguments)
merged.update(kwargs)
return merged
amount = calc_points(
scene_key,
is_member,
duration_minutes=duration_minutes,
extra_segments=extra_segments,
)
if amount <= 0:
logger.debug("points_gate: 免费场景 scene=%s,放行", scene_key)
yield None
return
def _execute_with_gate(
func: Callable,
args: tuple,
kwargs: dict,
scene_key: str,
per_unit: int | None,
unit_field: str | None,
quantity_field: str | None,
is_async: bool,
) -> Any:
"""积分扣费核心逻辑。"""
merged = _extract_kwargs(func, args, kwargs)
# 提取 current_user
current_user = merged.get("current_user")
if current_user is None:
# 尝试从位置参数中找
for arg in args:
if hasattr(arg, "user"):
current_user = arg
break
if not current_user:
raise HTTPException(status_code=401, detail="未登录")
# 提取 db session
db = merged.get("db")
if db is None:
raise HTTPException(status_code=500, detail="缺少数据库 session")
user = current_user.user
is_member = getattr(user, "is_member", False)
member_type = getattr(user, "member_type", None)
# ── 混剪场景:先检查免费额度 ──
if scene_key == "ai_video":
from packages.domain.points_service import PointsService
svc = PointsService()
if not is_member:
if svc.check_daily_free_clip(user.id, db):
svc.record_daily_free_clip(user.id, db)
kwargs["_points_deducted"] = 0
kwargs["_is_free_quota"] = True
if is_async:
return _run_async(func, args, kwargs)
return func(*args, **kwargs)
# ── 计算积分消耗 ──
if per_unit is not None:
total_points = per_unit
else:
from packages.domain.points_rules import calculate_points_cost
quantity = 1
duration = 0.0
request_body = merged.get("body") or merged.get("request") or merged.get("payload")
if request_body and unit_field:
duration = float(getattr(request_body, unit_field, 0) or 0)
if request_body and quantity_field:
quantity = int(getattr(request_body, quantity_field, 1) or 1)
total_points = calculate_points_cost(
result = points_svc.check_and_deduct(
user_id=user_id,
scene_key=scene_key,
duration_minutes=duration_minutes,
extra_segments=extra_segments,
description=desc,
is_member=is_member,
)
if not result.success:
logger.info(
"points_gate: 扣费失败 user=%s scene=%s reason=%s need=%d bal=%d",
user_id,
scene_key,
is_member,
quantity=quantity,
duration_minutes=duration,
member_type=member_type,
result.reason,
amount,
result.balance,
)
raise _insufficient_points(amount, result.balance, scene_name)
# 零消耗场景(如免费的声音克隆训练)直接放行
if total_points == 0:
kwargs["_points_deducted"] = 0
if is_async:
return _run_async(func, args, kwargs)
return func(*args, **kwargs)
# ── 扣减积分 ──
from packages.domain.points_service import PointsService
svc = PointsService()
job_id = merged.get("job_id", "") or ""
result = svc.deduct_points(user.id, total_points, scene_key, db, ref_id=str(job_id))
if not result["success"]:
raise HTTPException(
status_code=402,
detail={
"code": "INSUFFICIENT_POINTS",
"message": f"积分不足,需要 {total_points} 积分,当前余额 {result['balance']}",
"required": total_points,
"balance": result["balance"],
},
)
kwargs["_points_deducted"] = total_points
kwargs["_points_transaction_id"] = result["transaction_id"]
# ── 执行业务函数,失败则退还积分 ──
tx_id = result.transaction_id
logger.info(
"points_gate: 扣费成功 user=%s scene=%s amount=%d tx=%s",
user_id,
scene_key,
amount,
tx_id,
)
try:
if is_async:
return _run_async(func, args, kwargs)
return func(*args, **kwargs)
yield tx_id
except Exception:
svc.refund_points(user.id, total_points, scene_key, db, ref_id=str(job_id))
if tx_id:
try:
points_svc.refund(user_id, tx_id, reason=f"{scene_key} 业务失败: {desc}")
logger.info("points_gate: 业务异常已退款 user=%s tx=%s scene=%s", user_id, tx_id, scene_key)
except Exception:
logger.exception("points_gate: 退款失败 user=%s tx=%s", user_id, tx_id)
raise
def _run_async(func: Callable, args: tuple, kwargs: dict):
"""在 async wrapper 中 await 原始 async 函数。"""
return func(*args, **kwargs)