feat(#1895): 会员积分系统后端(P1-P4,默认关闭) #1923
@@ -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
@@ -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"],
|
||||
)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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 — 封面选定后正式入库成片库 ────────────────────
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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
@@ -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}
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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 {})
|
||||
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Executable → Regular
+45
-27
@@ -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,
|
||||
|
||||
@@ -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
|
||||
@@ -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 配置单例(统一入口)。"""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user