feat: 会员+积分系统后端 (#1895) #1919

Merged
auto-approve-bot merged 7 commits from feature/1895-membership-points into develop 2026-09-15 09:23:26 +08:00
27 changed files with 2981 additions and 0 deletions
+133
View File
@@ -0,0 +1,133 @@
"""add membership & points system
Revision ID: 076_membership_points
Revises: 075_add_sentence_timings
Create Date: 2026-09-15
"""
import sqlalchemy as sa
from sqlalchemy import text
from alembic import op
revision = "076_membership_points"
down_revision = "075_add_sentence_timings"
branch_labels = None
depends_on = None
def upgrade() -> None:
# 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")),
)
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")),
)
# 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("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()"),
),
)
# 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("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(),
nullable=False,
server_default=sa.text("NOW()"),
),
)
# 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("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("paid_at", sa.DateTime(), nullable=True),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
)
# 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("count", sa.Integer(), nullable=False, server_default=sa.text("0")),
sa.Column(
"updated_at",
sa.DateTime(),
nullable=False,
server_default=sa.text("NOW()"),
),
sa.UniqueConstraint(
"user_id",
"usage_date",
"usage_type",
name="uq_daily_usage_user_date_type",
),
)
def downgrade() -> None:
op.drop_table("daily_usage_records")
op.drop_table("points_orders")
op.drop_table("points_transactions")
op.drop_table("points_accounts")
with op.batch_alter_table("users") as batch:
batch.drop_column("points_balance")
batch.drop_column("member_expires_at")
batch.drop_column("member_type")
batch.drop_column("is_member")
+11
View File
@@ -17,6 +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.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
@@ -189,3 +190,13 @@ 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"],
)
+321
View File
@@ -0,0 +1,321 @@
"""积分 & 会员 API 路由 (#1895)
导出两个 router
- points_router: 积分相关路由,前缀 /points
- usage_router: 每日额度路由,前缀 /usage
"""
from __future__ import annotations
import logging
from datetime import datetime
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.points import (
DailyUsageResponse,
MembershipStatusResponse,
PointRuleItem,
PointsBalanceResponse,
PointsCheckRequest,
PointsCheckResponse,
PointsDeductRequest,
PointsOrderResponse,
PointsPackageItem,
PointsPackagesResponse,
PointsRechargeRequest,
PointsRefundRequest,
PointsRulesResponse,
PointsTransactionsResponse,
SimpleMessageResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session
from packages.domain.points_rules import (
FREE_USER_MULTIPLIER,
MEMBER_DISCOUNT,
POINTS_PACKAGES,
POINTS_SCENES,
calculate_points_cost,
)
from packages.domain.points_service import PointsService
logger = logging.getLogger(__name__)
# ── 两个 router ──
points_router = APIRouter()
usage_router = APIRouter()
def _get_service() -> PointsService:
return PointsService()
def _is_member(user: AuthenticatedUser) -> bool:
"""判断用户是否为付费会员。"""
return getattr(user.user, "is_member", False)
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(
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)
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),
)
@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),
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,
start_date=start_date,
end_date=end_date,
)
return PointsTransactionsResponse(**result)
@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():
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"),
)
)
return PointsRulesResponse(
rules=rules,
free_user_multiplier=FREE_USER_MULTIPLIER,
)
@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)
@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
+182
View File
@@ -0,0 +1,182 @@
"""积分 & 会员相关 Pydantic Schema (#1895)"""
from __future__ import annotations
from datetime import datetime
from typing import Any, 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="会员到期时间")
# ============ 流水 ============
class PointsTransactionItem(BaseModel):
"""单条积分流水"""
id: str
type: str = Field(..., description="类型: add/deduct")
source: str = Field(..., description="来源场景")
amount: int
balance_after: int
description: str = ""
ref_id: str = ""
created_at: Optional[str] = None
class PointsTransactionsResponse(BaseModel):
"""积分流水分页响应"""
items: list[PointsTransactionItem]
total: int
page: int
page_size: int
# ============ 规则 & 积分包 ============
class PointRuleItem(BaseModel):
"""单条积分规则"""
scene_key: str
name: str
base_points: int
unit: str
extra_per_30s: Optional[int] = None
class PointsRulesResponse(BaseModel):
"""所有积分消耗规则"""
rules: list[PointRuleItem]
free_user_multiplier: float = Field(..., description="免费用户积分上浮系数")
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
+72
View File
@@ -0,0 +1,72 @@
-- 会员 + 积分系统 (#1895)
-- users 表新增字段 + 4 张新表
-- 创建时间: 2026-09-11
-- 1. users 表新增字段
ALTER TABLE users ADD COLUMN IF NOT EXISTS is_member BOOLEAN NOT NULL DEFAULT FALSE;
ALTER TABLE users ADD COLUMN IF NOT EXISTS member_type VARCHAR(20);
ALTER TABLE users ADD COLUMN IF NOT EXISTS member_expires_at TIMESTAMP;
ALTER TABLE users ADD COLUMN IF NOT EXISTS points_balance INTEGER NOT NULL DEFAULT 0;
-- 2. points_accounts 积分账户表
CREATE TABLE IF NOT EXISTS points_accounts (
id VARCHAR(36) PRIMARY KEY,
user_id VARCHAR(36) NOT NULL UNIQUE REFERENCES users(id) ON DELETE CASCADE,
balance INTEGER NOT NULL DEFAULT 0,
total_earned INTEGER NOT NULL DEFAULT 0,
total_spent INTEGER NOT NULL DEFAULT 0,
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
updated_at TIMESTAMP NOT NULL DEFAULT NOW()
);
-- 3. points_transactions 积分流水表
CREATE TABLE IF NOT EXISTS points_transactions (
id VARCHAR(36) PRIMARY KEY,
user_id VARCHAR(36) NOT NULL REFERENCES users(id) ON DELETE CASCADE,
account_id VARCHAR(36) NOT NULL REFERENCES points_accounts(id) ON DELETE CASCADE,
type VARCHAR(20) NOT NULL,
source VARCHAR(50) NOT NULL,
amount INTEGER NOT NULL,
balance_after INTEGER NOT NULL,
description VARCHAR(255) DEFAULT '',
ref_id VARCHAR(100) DEFAULT '',
created_at TIMESTAMP NOT NULL DEFAULT NOW()
);
CREATE INDEX IF NOT EXISTS idx_points_tx_user ON points_transactions(user_id);
CREATE INDEX IF NOT EXISTS idx_points_tx_type ON points_transactions(type);
CREATE INDEX IF NOT EXISTS idx_points_tx_source ON points_transactions(source);
CREATE INDEX IF NOT EXISTS idx_points_tx_created ON points_transactions(created_at);
-- 4. points_orders 积分/会员订单表
CREATE TABLE IF NOT EXISTS points_orders (
id VARCHAR(36) PRIMARY KEY,
user_id VARCHAR(36) NOT NULL REFERENCES users(id) ON DELETE CASCADE,
order_type VARCHAR(20) NOT NULL,
product_code VARCHAR(50) NOT NULL,
amount_cents INTEGER NOT NULL,
original_amount_cents INTEGER NOT NULL DEFAULT 0,
discount REAL NOT NULL DEFAULT 1.0,
points_amount INTEGER NOT NULL DEFAULT 0,
status VARCHAR(20) NOT NULL DEFAULT 'pending',
payment_method VARCHAR(50),
payment_id VARCHAR(100),
paid_at TIMESTAMP,
created_at TIMESTAMP NOT NULL DEFAULT NOW()
);
CREATE INDEX IF NOT EXISTS idx_points_orders_user ON points_orders(user_id);
CREATE INDEX IF NOT EXISTS idx_points_orders_status ON points_orders(status);
-- 5. daily_usage_records 每日免费混剪计数
CREATE TABLE IF NOT EXISTS daily_usage_records (
id VARCHAR(36) PRIMARY KEY,
user_id VARCHAR(36) NOT NULL REFERENCES users(id) ON DELETE CASCADE,
usage_date DATE NOT NULL,
usage_type VARCHAR(50) NOT NULL DEFAULT 'free_clip',
count INTEGER NOT NULL DEFAULT 0,
updated_at TIMESTAMP NOT NULL DEFAULT NOW(),
UNIQUE(user_id, usage_date, usage_type)
);
CREATE INDEX IF NOT EXISTS idx_daily_usage_user_date ON daily_usage_records(user_id, usage_date);
@@ -0,0 +1,93 @@
from datetime import date, datetime, timezone
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import DailyUsageRecordModel
from packages.domain.daily_usage_record import DailyUsageRecord
class SQLAlchemyDailyUsageRepository:
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 = (
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:
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:
return record
model.count = record.count
model.updated_at = datetime.now(timezone.utc)
self.session.add(model)
self.session.commit()
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
model.count += 1
model.updated_at = datetime.now(timezone.utc)
self.session.add(model)
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,
)
@@ -39,6 +39,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)
member_expires_at = Column(DateTime, nullable=True)
points_balance = Column(Integer, nullable=False, default=0)
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -743,3 +748,68 @@ class AiAvatarRenderJob(Base):
completed_at = Column(DateTime, nullable=True)
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 PointsAccountModel(Base):
"""积分账户 ORM 模型 (#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))
class PointsTransactionModel(Base):
"""积分流水 ORM 模型 (#1895)"""
__tablename__ = "points_transactions"
id = Column(String(36), primary_key=True)
user_id = Column(String(36), nullable=False, index=True)
account_id = Column(String(36), nullable=False, index=True)
type = Column(String(20), nullable=False, index=True) # earn / spend / refund
source = Column(String(50), nullable=False, index=True)
amount = Column(Integer, nullable=False)
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))
class PointsOrderModel(Base):
"""积分/会员订单 ORM 模型 (#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)
discount = Column(Float, nullable=False, default=1.0)
points_amount = Column(Integer, nullable=False, default=0)
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))
class DailyUsageRecordModel(Base):
"""每日使用记录 ORM 模型 (#1895)"""
__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))
@@ -0,0 +1,55 @@
from datetime import datetime, timezone
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:
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,
)
self.session.add(model)
self.session.commit()
return account
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 update_balance(self, account: PointsAccount) -> PointsAccount:
model = self.session.query(PointsAccountModel).filter(PointsAccountModel.id == account.id).first()
if model is None:
return account
model.balance = account.balance
model.total_earned = account.total_earned
model.total_spent = account.total_spent
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,
)
@@ -0,0 +1,96 @@
from datetime import datetime
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:
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,
)
self.session.add(model)
self.session.commit()
return order
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 update_status(
self,
order_id: 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()
if model is None:
return None
model.status = status
if payment_id is not None:
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)
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,
)
@@ -0,0 +1,65 @@
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:
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,
)
self.session.add(model)
self.session.commit()
return transaction
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)
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,
)
+8
View File
@@ -6,6 +6,7 @@ from .classification import (
ClassificationJobStatus,
)
from .cover_template import CoverTemplate
from .daily_usage_record import DailyUsageRecord
from .duplication import DuplicateSegment, DuplicationRecord
from .edit_plan import EditPlan, EditPlanStatus
from .edit_plan_clip import EditPlanClip, EditPlanClipStatus
@@ -25,6 +26,9 @@ from .entities import (
from .generated_video import GeneratedVideo
from .generation_task import GenerationTask, GenerationTaskStatus
from .job import Job, JobStatus, JobType
from .points_account import PointsAccount
from .points_order import PointsOrder
from .points_transaction import PointsTransaction
from .smart_match import SmartMatchResult, score_asset, smart_select_assets
from .tag import Tag
from .template_clip_config import ClipType, TemplateClipConfig, TransitionEffect
@@ -34,6 +38,10 @@ from .voice_library import VoiceLibraryItem
__all__ = [
"Asset",
"AssetClassification",
"DailyUsageRecord",
"PointsAccount",
"PointsOrder",
"PointsTransaction",
"AssetLibrary",
"AssetLibraryKind",
"AssetStatus",
+29
View File
@@ -0,0 +1,29 @@
from __future__ import annotations
from dataclasses import dataclass, field
from datetime import date, datetime, timezone
from uuid import uuid4
@dataclass(slots=True)
class DailyUsageRecord:
id: str
user_id: str
usage_date: date
usage_type: str = "free_clip"
count: int = 0
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@classmethod
def create(
cls,
user_id: str,
usage_date: date,
usage_type: str = "free_clip",
) -> "DailyUsageRecord":
return cls(
id=uuid4().hex,
user_id=user_id,
usage_date=usage_date,
usage_type=usage_type,
)
+20
View File
@@ -0,0 +1,20 @@
from __future__ import annotations
from dataclasses import dataclass, field
from datetime import datetime, timezone
from uuid import uuid4
@dataclass(slots=True)
class PointsAccount:
id: str
user_id: str
balance: int = 0
total_earned: int = 0
total_spent: int = 0
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@classmethod
def create(cls, user_id: str) -> "PointsAccount":
return cls(id=uuid4().hex, user_id=user_id)
+45
View File
@@ -0,0 +1,45 @@
from __future__ import annotations
from dataclasses import dataclass, field
from datetime import datetime, timezone
from uuid import uuid4
@dataclass(slots=True)
class PointsOrder:
id: str
user_id: str
order_type: str # membership / points
product_code: str
amount_cents: int
original_amount_cents: int = 0
discount: float = 1.0
points_amount: int = 0
status: str = "pending"
payment_method: str | None = None
payment_id: str | None = None
paid_at: datetime | None = None
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@classmethod
def create(
cls,
user_id: str,
order_type: str,
product_code: str,
amount_cents: int,
*,
original_amount_cents: int = 0,
discount: float = 1.0,
points_amount: int = 0,
) -> "PointsOrder":
return cls(
id=uuid4().hex,
user_id=user_id,
order_type=order_type,
product_code=product_code,
amount_cents=amount_cents,
original_amount_cents=original_amount_cents,
discount=discount,
points_amount=points_amount,
)
+107
View File
@@ -0,0 +1,107 @@
"""积分消耗规则配置 (#1895)"""
from __future__ import annotations
import math
# ============ 场景定义 ============
# 每个场景: base_points(基础积分), unit(计费单位), name(显示名称)
POINTS_SCENES: dict[str, dict] = {
"ai_voice": {"base_points": 1, "unit": "分钟", "name": "AI 配音"},
"ai_video": {
"base_points": 3,
"unit": "",
"name": "智能混剪",
"extra_per_30s": 1,
},
"ai_digital_human": {"base_points": 15, "unit": "分钟", "name": "AI 数字人"},
"voice_clone_train": {"base_points": 0, "unit": "", "name": "声音克隆训练"},
"voice_clone_synth": {"base_points": 1, "unit": "分钟", "name": "声音克隆合成"},
"douyin_extract": {"base_points": 1, "unit": "", "name": "抖音链接提取"},
"ai_rewrite": {"base_points": 1, "unit": "", "name": "AI 改写文案"},
"ai_title": {"base_points": 1, "unit": "", "name": "AI 标题生成"},
"ai_cover": {"base_points": 1, "unit": "", "name": "AI 封面生成"},
}
# 免费用户积分消耗上浮系数
FREE_USER_MULTIPLIER = 1.15
# ============ 积分包定义 ============
POINTS_PACKAGES: dict[str, dict] = {
"starter_pack": {"name": "体验包", "points": 100, "price_cents": 990},
"basic_pack": {"name": "基础包", "points": 500, "price_cents": 3900},
"pro_pack": {"name": "专业包", "points": 2000, "price_cents": 12900},
}
# ============ 会员定价 ============
MEMBERSHIP_PRICES: dict[str, dict] = {
"monthly": {"name": "月卡", "price_cents": 1990, "duration_days": 30},
"quarterly": {"name": "季卡", "price_cents": 3990, "duration_days": 90},
"yearly": {"name": "年卡", "price_cents": 15900, "duration_days": 365},
}
# 会员积分折扣(付费会员按此系数打折)
MEMBER_DISCOUNT: dict[str, float] = {
"monthly": 0.9,
"quarterly": 0.87,
"yearly": 0.8,
}
# 每日免费混剪次数(免费用户)
DAILY_FREE_CLIP_LIMIT = 2
def calculate_points_cost(
scene_key: str,
is_member: bool,
quantity: int = 1,
duration_minutes: float = 0,
member_type: str | None = None,
) -> int:
"""计算指定场景的积分消耗。
Args:
scene_key: 场景标识,如 "ai_voice""ai_video"
is_member: 是否付费会员
quantity: 数量(按次计费场景)
duration_minutes: 时长分钟数(按时长计费场景)
member_type: 会员类型 (monthly/quarterly/yearly),用于折扣
Returns:
实际消耗积分(已含免费用户 ×1.15 上浮或会员折扣)
Raises:
ValueError: 未知场景标识
"""
scene = POINTS_SCENES.get(scene_key)
if not scene:
raise ValueError(f"Unknown points scene: {scene_key}")
base = scene["base_points"]
if base == 0:
return 0
# —— 计算基础消耗 ——
unit = scene["unit"]
if unit == "分钟":
total_base = base * max(1, math.ceil(duration_minutes))
elif unit in ("", "", ""):
total_base = base * quantity
# 混剪特殊逻辑:视频超过 30s 后每 +30s 额外加 1 积分
if scene_key == "ai_video" and duration_minutes > 0.5:
extra_segments = math.ceil((duration_minutes * 60 - 30) / 30)
if extra_segments > 0:
total_base += scene.get("extra_per_30s", 1) * extra_segments
else:
total_base = base
# —— 会员折扣 / 免费用户上浮 ——
if is_member and member_type and member_type in MEMBER_DISCOUNT:
total_base = max(1, math.floor(total_base * MEMBER_DISCOUNT[member_type]))
elif not is_member:
total_base = math.ceil(total_base * FREE_USER_MULTIPLIER)
return total_base
+571
View File
@@ -0,0 +1,571 @@
"""积分服务层 — 积分账户、扣减、充值、流水、每日免费额度 (#1895)
直接操作 SQLAlchemy session,不走 Repository 抽象层,简化事务处理。
"""
from __future__ import annotations
import logging
import uuid
from datetime import datetime, timedelta, timezone
from typing import Any
from sqlalchemy.orm import Session
from packages.domain.points_rules import (
DAILY_FREE_CLIP_LIMIT,
POINTS_PACKAGES,
)
logger = logging.getLogger(__name__)
# ── 延迟导入模型(避免循环/顺序依赖) ────────────────────────────────────
def _get_models():
"""延迟获取积分相关模型类。"""
from packages.adapters.sqlalchemy_impl.models import (
DailyUsageRecordModel,
PointsAccountModel,
PointsOrderModel,
PointsTransactionModel,
UserModel,
)
return (
PointsAccountModel,
PointsTransactionModel,
PointsOrderModel,
DailyUsageRecordModel,
UserModel,
)
def _get_redis_client():
"""获取 Redis 客户端,用于每日额度缓存。"""
try:
import redis as redis_lib
from app.config import settings
return redis_lib.from_url(settings.REDIS_URL, decode_responses=True)
except Exception:
return None
class PointsService:
"""积分核心服务。"""
# ──────────────── 账户管理 ────────────────
def get_or_create_account(self, user_id: str, db: Session) -> dict[str, Any]:
"""获取或创建积分账户,返回账户快照。"""
PointsAccountModel, _, _, _, _ = _get_models()
account = db.query(PointsAccountModel).filter(PointsAccountModel.user_id == user_id).first()
if account is None:
account = PointsAccountModel(
id=uuid.uuid4().hex,
user_id=user_id,
balance=0,
total_earned=0,
total_spent=0,
)
db.add(account)
db.flush()
return {
"id": account.id,
"user_id": account.user_id,
"balance": account.balance,
"total_earned": account.total_earned,
"total_spent": account.total_spent,
}
# ──────────────── 余额检查 ────────────────
def check_balance(self, user_id: str, amount: int, db: Session) -> dict[str, Any]:
"""检查余额是否足够。"""
account_data = self.get_or_create_account(user_id, db)
balance = account_data["balance"]
return {
"sufficient": balance >= amount,
"balance": balance,
"required": amount,
"remaining_after": balance - amount,
}
# ──────────────── 积分扣减(事务性) ────────────────
def deduct_points(
self,
user_id: str,
amount: int,
source: str,
db: Session,
description: str = "",
ref_id: str = "",
) -> dict[str, Any]:
"""扣减积分(事务性:SELECT FOR UPDATE → 检查余额 → 扣减 → 流水 → 同步用户表)。
Returns:
{"success": True/False, "balance": int, "transaction_id": str|None}
"""
PointsAccountModel, PointsTransactionModel, _, _, UserModel = _get_models()
try:
# 1. 行锁获取账户
account = (
db.query(PointsAccountModel).filter(PointsAccountModel.user_id == user_id).with_for_update().first()
)
if account is None:
account = PointsAccountModel(
id=uuid.uuid4().hex,
user_id=user_id,
balance=0,
total_earned=0,
total_spent=0,
)
db.add(account)
db.flush()
# 2. 检查余额
if account.balance < amount:
return {
"success": False,
"balance": account.balance,
"transaction_id": None,
}
# 3. 扣减余额
account.balance -= amount
account.total_spent += amount
# 4. 创建流水
txn_id = uuid.uuid4().hex
txn = PointsTransactionModel(
id=txn_id,
user_id=user_id,
account_id=account.id,
type="deduct",
source=source,
amount=amount,
balance_after=account.balance,
description=description or f"积分扣减: {source}",
ref_id=ref_id,
)
db.add(txn)
# 5. 同步用户表 points_balance
db.execute(
UserModel.__table__.update()
.where(UserModel.__table__.c.id == user_id)
.values(points_balance=account.balance)
)
db.commit()
return {
"success": True,
"balance": account.balance,
"transaction_id": txn_id,
}
except Exception:
db.rollback()
logger.exception(
"积分扣减失败: user_id=%s, amount=%d, source=%s",
user_id,
amount,
source,
)
return {"success": False, "balance": 0, "transaction_id": None}
# ──────────────── 积分增加 ────────────────
def add_points(
self,
user_id: str,
amount: int,
source: str,
db: Session,
description: str = "",
ref_id: str = "",
) -> dict[str, Any]:
"""增加积分(充值/赠送/退款)。"""
PointsAccountModel, PointsTransactionModel, _, _, UserModel = _get_models()
try:
account = (
db.query(PointsAccountModel).filter(PointsAccountModel.user_id == user_id).with_for_update().first()
)
if account is None:
account = PointsAccountModel(
id=uuid.uuid4().hex,
user_id=user_id,
balance=0,
total_earned=0,
total_spent=0,
)
db.add(account)
db.flush()
account.balance += amount
account.total_earned += amount
txn_id = uuid.uuid4().hex
txn = PointsTransactionModel(
id=txn_id,
user_id=user_id,
account_id=account.id,
type="add",
source=source,
amount=amount,
balance_after=account.balance,
description=description or f"积分增加: {source}",
ref_id=ref_id,
)
db.add(txn)
db.execute(
UserModel.__table__.update()
.where(UserModel.__table__.c.id == user_id)
.values(points_balance=account.balance)
)
db.commit()
return {
"success": True,
"balance": account.balance,
"transaction_id": txn_id,
}
except Exception:
db.rollback()
logger.exception(
"积分增加失败: user_id=%s, amount=%d, source=%s",
user_id,
amount,
source,
)
return {"success": False, "balance": 0, "transaction_id": None}
# ──────────────── 积分退还 ────────────────
def refund_points(
self,
user_id: str,
amount: int,
source: str,
db: Session,
ref_id: str = "",
description: str = "",
) -> dict[str, Any]:
"""退还积分(业务失败回退)。内部调用 add_pointssource 前缀 refund:。"""
return self.add_points(
user_id=user_id,
amount=amount,
source=f"refund:{source}",
db=db,
description=description or f"积分退还: {source}",
ref_id=ref_id,
)
# ──────────────── 流水查询 ────────────────
def get_transactions(
self,
user_id: str,
db: Session,
page: int = 1,
page_size: int = 20,
type_filter: str | None = None,
source_filter: str | None = None,
start_date: datetime | None = None,
end_date: datetime | None = None,
) -> dict[str, Any]:
"""查询积分流水(分页+筛选)。"""
_, PointsTransactionModel, _, _, _ = _get_models()
query = db.query(PointsTransactionModel).filter(PointsTransactionModel.user_id == user_id)
if type_filter:
query = query.filter(PointsTransactionModel.type == type_filter)
if source_filter:
query = query.filter(PointsTransactionModel.source == source_filter)
if start_date:
query = query.filter(PointsTransactionModel.created_at >= start_date)
if end_date:
query = query.filter(PointsTransactionModel.created_at <= end_date)
total = query.count()
items = (
query.order_by(PointsTransactionModel.created_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
.all()
)
return {
"items": [
{
"id": item.id,
"type": item.type,
"source": item.source,
"amount": item.amount,
"balance_after": item.balance_after,
"description": item.description,
"ref_id": item.ref_id,
"created_at": (item.created_at.isoformat() if item.created_at else None),
}
for item in items
],
"total": total,
"page": page,
"page_size": page_size,
}
# ──────────────── 每日免费混剪额度 ────────────────
def _daily_key(self, user_id: str) -> str:
"""生成 Redis 每日额度 key。格式: daily_usage:{user_id}:{YYYYMMDD}:free_clip"""
today = datetime.now(timezone.utc).strftime("%Y%m%d")
return f"daily_usage:{user_id}:{today}:free_clip"
def check_daily_free_clip(self, user_id: str, db: Session) -> bool:
"""检查今日是否还有免费混剪额度。
优先查 Redis,Redis 不可用时降级到 DB。
"""
redis_client = _get_redis_client()
if redis_client:
try:
key = self._daily_key(user_id)
current = redis_client.get(key)
if current is None:
return True
return int(current) < DAILY_FREE_CLIP_LIMIT
except Exception:
logger.warning("Redis 不可用,降级到 DB 查询每日额度")
# 降级到 DB
_, _, _, DailyUsageRecordModel, _ = _get_models()
today_start = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
record = (
db.query(DailyUsageRecordModel)
.filter(
DailyUsageRecordModel.user_id == user_id,
DailyUsageRecordModel.usage_type == "free_clip",
DailyUsageRecordModel.usage_date >= today_start,
)
.first()
)
if record is None:
return True
return record.count < DAILY_FREE_CLIP_LIMIT
def record_daily_free_clip(self, user_id: str, db: Session) -> bool:
"""记录使用一次免费混剪。
先 INCR Redis;如果超限回退 Redis。DB 使用 upsert 语义(唯一约束)。
"""
redis_client = _get_redis_client()
if redis_client:
try:
key = self._daily_key(user_id)
new_count = redis_client.incr(key)
if new_count == 1:
redis_client.expire(key, 48 * 3600) # TTL 48h
if new_count <= DAILY_FREE_CLIP_LIMIT:
return True
# 超限,回退 Redis
redis_client.decr(key)
except Exception:
logger.warning("Redis 不可用,降级到 DB 记录每日额度")
# 降级/兜底到 DBupsert 语义)
_, _, _, DailyUsageRecordModel, _ = _get_models()
today_start = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
record = (
db.query(DailyUsageRecordModel)
.filter(
DailyUsageRecordModel.user_id == user_id,
DailyUsageRecordModel.usage_type == "free_clip",
DailyUsageRecordModel.usage_date >= today_start,
)
.first()
)
if record is None:
if DAILY_FREE_CLIP_LIMIT <= 0:
return False
record = DailyUsageRecordModel(
id=uuid.uuid4().hex,
user_id=user_id,
usage_type="free_clip",
usage_date=datetime.now(timezone.utc),
count=1,
)
db.add(record)
else:
if record.count >= DAILY_FREE_CLIP_LIMIT:
return False
record.count += 1
db.commit()
return True
def get_daily_usage(self, user_id: str, db: Session) -> dict[str, Any]:
"""查询今日免费额度使用情况。"""
redis_client = _get_redis_client()
used = 0
if redis_client:
try:
key = self._daily_key(user_id)
val = redis_client.get(key)
used = int(val) if val else 0
except Exception:
pass
if used == 0:
# 从 DB 查
_, _, _, DailyUsageRecordModel, _ = _get_models()
today_start = datetime.now(timezone.utc).replace(hour=0, minute=0, second=0, microsecond=0)
record = (
db.query(DailyUsageRecordModel)
.filter(
DailyUsageRecordModel.user_id == user_id,
DailyUsageRecordModel.usage_type == "free_clip",
DailyUsageRecordModel.usage_date >= today_start,
)
.first()
)
used = record.count if record else 0
now = datetime.now(timezone.utc)
tomorrow = (now + timedelta(days=1)).replace(hour=0, minute=0, second=0, microsecond=0)
return {
"free_clips_used": used,
"free_clips_limit": DAILY_FREE_CLIP_LIMIT,
"free_clips_remaining": max(0, DAILY_FREE_CLIP_LIMIT - used),
"reset_at": tomorrow.isoformat(),
}
# ──────────────── 订单管理 ────────────────
def create_order(
self,
user_id: str,
order_type: str,
product_code: str,
db: Session,
) -> dict[str, Any]:
"""创建积分充值或会员购买订单。
Args:
order_type: "points""membership"
product_code: 积分包 code (如 "starter_pack") 或会员类型 (如 "monthly")
"""
_, _, PointsOrderModel, _, _ = _get_models()
amount_cents = 0
points_amount = 0
if order_type == "points":
package = POINTS_PACKAGES.get(product_code)
if not package:
raise ValueError(f"Unknown points package: {product_code}")
amount_cents = package["price_cents"]
points_amount = package["points"]
elif order_type == "membership":
from packages.domain.points_rules import MEMBERSHIP_PRICES
membership = MEMBERSHIP_PRICES.get(product_code)
if not membership:
raise ValueError(f"Unknown membership type: {product_code}")
amount_cents = membership["price_cents"]
else:
raise ValueError(f"Unknown order type: {order_type}")
order = PointsOrderModel(
id=uuid.uuid4().hex,
user_id=user_id,
order_type=order_type,
product_code=product_code,
amount_cents=amount_cents,
original_amount_cents=amount_cents,
points_amount=points_amount,
status="pending",
)
db.add(order)
db.commit()
return {
"id": order.id,
"order_type": order.order_type,
"product_code": order.product_code,
"amount_cents": order.amount_cents,
"status": order.status,
"created_at": (order.created_at.isoformat() if order.created_at else None),
}
def confirm_payment(
self,
order_id: str,
payment_id: str,
db: Session,
) -> dict[str, Any]:
"""确认支付 → 更新订单状态 → 发放积分或会员。"""
_, _, PointsOrderModel, _, UserModel = _get_models()
try:
order = db.query(PointsOrderModel).filter(PointsOrderModel.id == order_id).with_for_update().first()
if order is None:
return {"success": False, "message": "订单不存在"}
if order.status != "pending":
return {"success": False, "message": f"订单状态异常: {order.status}"}
# 更新订单状态
order.status = "paid"
order.payment_id = payment_id
order.paid_at = datetime.now(timezone.utc)
if order.order_type == "points":
# 发放积分
self.add_points(
user_id=order.user_id,
amount=order.points_amount,
source=f"recharge:{order.product_code}",
db=db,
description=f"积分充值: {order.product_code}",
ref_id=order.id,
)
elif order.order_type == "membership":
# 激活会员
from packages.domain.points_rules import MEMBERSHIP_PRICES
membership = MEMBERSHIP_PRICES.get(order.product_code, {})
duration_days = membership.get("duration_days", 30)
user = db.query(UserModel).filter(UserModel.id == order.user_id).first()
if user:
now = datetime.now(timezone.utc)
current_expires = user.member_expires_at or now
if current_expires < now:
current_expires = now
user.member_expires_at = current_expires + timedelta(days=duration_days)
user.member_type = order.product_code
user.is_member = True
db.commit()
return {
"success": True,
"message": "支付确认成功",
"order_id": order_id,
}
except Exception:
db.rollback()
logger.exception("确认支付失败: order_id=%s", order_id)
return {"success": False, "message": "确认支付异常"}
+43
View File
@@ -0,0 +1,43 @@
from __future__ import annotations
from dataclasses import dataclass, field
from datetime import datetime, timezone
from uuid import uuid4
@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 = field(default_factory=lambda: datetime.now(timezone.utc))
@classmethod
def create(
cls,
user_id: str,
account_id: str,
type: str,
source: str,
amount: int,
balance_after: int,
description: str = "",
ref_id: str = "",
) -> "PointsTransaction":
return cls(
id=uuid4().hex,
user_id=user_id,
account_id=account_id,
type=type,
source=source,
amount=amount,
balance_after=balance_after,
description=description,
ref_id=ref_id,
)
View File
+181
View File
@@ -0,0 +1,181 @@
"""AI 功能入口的积分扣费装饰器 (#1895)
支持 sync 和 async 函数。业务失败时自动退还积分。
"""
from __future__ import annotations
import asyncio
import functools
import inspect
import logging
from typing import Any, Callable
from fastapi import HTTPException
logger = logging.getLogger(__name__)
def points_gate(
scene_key: str,
per_unit: int | None = None,
unit_field: str | None = None,
quantity_field: str | None = None,
) -> Callable:
"""AI 功能入口积分扣费装饰器。
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)):
...
"""
def decorator(func: Callable) -> Callable:
is_async = asyncio.iscoroutinefunction(func)
@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
)
@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 is_async:
return async_wrapper
return sync_wrapper
return decorator
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
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(
scene_key,
is_member,
quantity=quantity,
duration_minutes=duration,
member_type=member_type,
)
# 零消耗场景(如免费的声音克隆训练)直接放行
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"]
# ── 执行业务函数,失败则退还积分 ──
try:
if is_async:
return _run_async(func, args, kwargs)
return func(*args, **kwargs)
except Exception:
svc.refund_points(user.id, total_points, scene_key, db, ref_id=str(job_id))
raise
def _run_async(func: Callable, args: tuple, kwargs: dict):
"""在 async wrapper 中 await 原始 async 函数。"""
return func(*args, **kwargs)
+20
View File
@@ -0,0 +1,20 @@
from abc import ABC, abstractmethod
from datetime import date
from packages.domain.daily_usage_record import DailyUsageRecord
class DailyUsageRepository(ABC):
@abstractmethod
def create(self, record: DailyUsageRecord) -> DailyUsageRecord: ...
@abstractmethod
def get_by_user_and_date(
self, user_id: str, usage_date: date, usage_type: str = "free_clip"
) -> DailyUsageRecord | None: ...
@abstractmethod
def update_count(self, record: DailyUsageRecord) -> DailyUsageRecord: ...
@abstractmethod
def upsert(self, user_id: str, usage_date: date, usage_type: str = "free_clip") -> DailyUsageRecord: ...
@@ -0,0 +1,14 @@
from abc import ABC, abstractmethod
from packages.domain.points_account import PointsAccount
class PointsAccountRepository(ABC):
@abstractmethod
def create(self, account: PointsAccount) -> PointsAccount: ...
@abstractmethod
def get_by_user_id(self, user_id: str) -> PointsAccount | None: ...
@abstractmethod
def update_balance(self, account: PointsAccount) -> PointsAccount: ...
+33
View File
@@ -0,0 +1,33 @@
from abc import ABC, abstractmethod
from datetime import datetime
from packages.domain.points_order import PointsOrder
class PointsOrderRepository(ABC):
@abstractmethod
def create(self, order: PointsOrder) -> PointsOrder: ...
@abstractmethod
def get(self, order_id: str) -> PointsOrder | None: ...
@abstractmethod
def update_status(
self,
order_id: str,
status: str,
*,
payment_id: str | None = None,
paid_at: datetime | None = None,
) -> PointsOrder | None: ...
@abstractmethod
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]: ...
@@ -0,0 +1,19 @@
from abc import ABC, abstractmethod
from packages.domain.points_transaction import PointsTransaction
class PointsTransactionRepository(ABC):
@abstractmethod
def create(self, transaction: PointsTransaction) -> PointsTransaction: ...
@abstractmethod
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]: ...
+173
View File
@@ -0,0 +1,173 @@
"""points_gate 中间件单元测试 (#1895)"""
from __future__ import annotations
import asyncio
from unittest.mock import MagicMock, patch
import pytest
from fastapi import HTTPException
from packages.middleware.points_gate import _execute_with_gate, _extract_kwargs, points_gate
def _make_user(user_id="user-1", is_member=False, member_type=None):
user = MagicMock()
user.id = user_id
user.is_member = is_member
user.member_type = member_type
return user
def _make_current_user(user_id="user-1", is_member=False, member_type=None):
cu = MagicMock()
cu.user = _make_user(user_id, is_member, member_type)
return cu
class TestExtractKwargs:
def test_basic_extraction(self):
def fn(a, b, c=None):
pass
result = _extract_kwargs(fn, (1, 2), {"c": 3})
assert result == {"a": 1, "b": 2, "c": 3}
class TestPointsGateSync:
def test_no_user_raises_401(self):
@points_gate("ai_rewrite")
def my_func(db=None):
return "ok"
with pytest.raises(HTTPException) as exc_info:
my_func(db=MagicMock())
assert exc_info.value.status_code == 401
def test_no_db_raises_500(self):
@points_gate("ai_rewrite")
def my_func(current_user=None, db=None):
return "ok"
with pytest.raises(HTTPException) as exc_info:
my_func(current_user=_make_current_user(), db=None)
assert exc_info.value.status_code == 500
def test_zero_cost_scene_passes_through(self):
@points_gate("voice_clone_train")
def my_func(current_user=None, db=None, **kwargs):
return kwargs.get("_points_deducted", -1)
mock_db = MagicMock()
cu = _make_current_user()
result = my_func(current_user=cu, db=mock_db)
assert result == 0
class TestPointsGateExecuteLogic:
def test_insufficient_points_raises_402(self):
cu = _make_current_user()
db = MagicMock()
mock_svc = MagicMock()
mock_svc.deduct_points.return_value = {"success": False, "balance": 2, "transaction_id": None}
def my_func(current_user=cu, db=db, **kwargs):
return "ok"
with patch("packages.domain.points_service.PointsService", return_value=mock_svc):
with pytest.raises(HTTPException) as exc_info:
_execute_with_gate(
my_func, (), {"current_user": cu, "db": db}, "ai_rewrite", None, None, None, is_async=False
)
assert exc_info.value.status_code == 402
def test_free_scene_passes_through(self):
cu = _make_current_user()
db = MagicMock()
def my_func(current_user=cu, db=db, **kwargs):
return "result"
result = _execute_with_gate(
my_func, (), {"current_user": cu, "db": db}, "voice_clone_train", None, None, None, is_async=False
)
assert result == "result"
def test_per_unit_fixed_cost(self):
cu = _make_current_user()
db = MagicMock()
mock_svc = MagicMock()
mock_svc.deduct_points.return_value = {"success": True, "balance": 90, "transaction_id": "t1"}
def my_func(current_user=cu, db=db, **kwargs):
return kwargs.get("_points_deducted", 0)
with patch("packages.domain.points_service.PointsService", return_value=mock_svc):
result = _execute_with_gate(
my_func,
(),
{"current_user": cu, "db": db},
"ai_rewrite",
per_unit=10,
unit_field=None,
quantity_field=None,
is_async=False,
)
assert result == 10
mock_svc.deduct_points.assert_called_once()
def test_refund_on_failure(self):
cu = _make_current_user()
db = MagicMock()
mock_svc = MagicMock()
mock_svc.deduct_points.return_value = {"success": True, "balance": 90, "transaction_id": "t1"}
def failing_func(current_user=cu, db=db, **kwargs):
raise RuntimeError("business error")
with patch("packages.domain.points_service.PointsService", return_value=mock_svc):
with pytest.raises(RuntimeError, match="business error"):
_execute_with_gate(
failing_func,
(),
{"current_user": cu, "db": db},
"ai_rewrite",
per_unit=10,
unit_field=None,
quantity_field=None,
is_async=False,
)
mock_svc.refund_points.assert_called_once()
def test_ai_video_free_quota_for_free_user(self):
cu = _make_current_user(is_member=False)
db = MagicMock()
mock_svc = MagicMock()
mock_svc.check_daily_free_clip.return_value = True
mock_svc.record_daily_free_clip.return_value = True
def my_func(current_user=cu, db=db, **kwargs):
return kwargs.get("_is_free_quota", False)
with patch("packages.domain.points_service.PointsService", return_value=mock_svc):
result = _execute_with_gate(
my_func, (), {"current_user": cu, "db": db}, "ai_video", None, None, None, is_async=False
)
assert result is True
class TestPointsGateAsync:
@pytest.mark.asyncio
async def test_async_func_supported(self):
cu = _make_current_user()
db = MagicMock()
mock_svc = MagicMock()
mock_svc.deduct_points.return_value = {"success": True, "balance": 90, "transaction_id": "t1"}
@points_gate("ai_rewrite", per_unit=5)
async def my_async_func(current_user=None, db=None, **kwargs):
return kwargs.get("_points_deducted", 0)
with patch("packages.domain.points_service.PointsService", return_value=mock_svc):
result = await my_async_func(current_user=cu, db=db)
assert result == 5
+293
View File
@@ -0,0 +1,293 @@
"""积分系统 Repository 层单元测试 (#1895)"""
from __future__ import annotations
import uuid
from datetime import date, datetime, timezone
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from packages.adapters.sqlalchemy_impl.daily_usage_repository import SQLAlchemyDailyUsageRepository
from packages.adapters.sqlalchemy_impl.models import (
Base,
DailyUsageRecordModel,
PointsAccountModel,
PointsOrderModel,
PointsTransactionModel,
UserModel,
)
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.daily_usage_record import DailyUsageRecord
from packages.domain.points_account import PointsAccount
from packages.domain.points_order import PointsOrder
from packages.domain.points_transaction import PointsTransaction
@pytest.fixture()
def db_session():
engine = create_engine("sqlite://", echo=False)
Base.metadata.create_all(engine)
SessionLocal = sessionmaker(bind=engine)
session = SessionLocal()
# 创建测试用户
user = UserModel(
id="test-user-1",
email="test@example.com",
username="testuser",
display_name="Test User",
password_hash="xxx",
)
session.add(user)
session.commit()
yield session
session.close()
@pytest.fixture()
def account_repo(db_session):
return SQLAlchemyPointsAccountRepository(db_session)
@pytest.fixture()
def txn_repo(db_session):
return SQLAlchemyPointsTransactionRepository(db_session)
@pytest.fixture()
def order_repo(db_session):
return SQLAlchemyPointsOrderRepository(db_session)
@pytest.fixture()
def daily_repo(db_session):
return SQLAlchemyDailyUsageRepository(db_session)
# ── PointsAccountRepository ──
class TestPointsAccountRepository:
def test_create_and_get(self, account_repo):
account = PointsAccount.create(user_id="test-user-1")
result = account_repo.create(account)
assert result.user_id == "test-user-1"
fetched = account_repo.get_by_user_id("test-user-1")
assert fetched is not None
assert fetched.id == account.id
assert fetched.balance == 0
def test_get_nonexistent(self, account_repo):
result = account_repo.get_by_user_id("nonexistent")
assert result is None
def test_update_balance(self, account_repo):
account = PointsAccount.create(user_id="test-user-1")
account_repo.create(account)
account.balance = 100
account.total_earned = 150
account.total_spent = 50
updated = account_repo.update_balance(account)
assert updated.balance == 100
fetched = account_repo.get_by_user_id("test-user-1")
assert fetched.balance == 100
assert fetched.total_earned == 150
assert fetched.total_spent == 50
# ── PointsTransactionRepository ──
class TestPointsTransactionRepository:
def test_create_and_list(self, txn_repo):
txn = PointsTransaction.create(
user_id="test-user-1",
account_id="acc-1",
type="add",
source="recharge",
amount=100,
balance_after=100,
description="充值",
)
txn_repo.create(txn)
items, total = txn_repo.list_by_user("test-user-1")
assert total == 1
assert items[0].amount == 100
assert items[0].type == "add"
def test_list_with_type_filter(self, txn_repo):
for t in ["add", "deduct", "add"]:
txn = PointsTransaction.create(
user_id="test-user-1",
account_id="acc-1",
type=t,
source="test",
amount=10,
balance_after=10,
)
txn_repo.create(txn)
items, total = txn_repo.list_by_user("test-user-1", type="add")
assert total == 2
def test_list_with_source_filter(self, txn_repo):
for s in ["recharge", "ai_voice", "recharge"]:
txn = PointsTransaction.create(
user_id="test-user-1",
account_id="acc-1",
type="add",
source=s,
amount=10,
balance_after=10,
)
txn_repo.create(txn)
items, total = txn_repo.list_by_user("test-user-1", source="recharge")
assert total == 2
def test_list_pagination(self, txn_repo):
for _i in range(5):
txn = PointsTransaction.create(
user_id="test-user-1",
account_id="acc-1",
type="add",
source="test",
amount=10,
balance_after=10,
)
txn_repo.create(txn)
items, total = txn_repo.list_by_user("test-user-1", page=1, page_size=3)
assert total == 5
assert len(items) == 3
items2, _ = txn_repo.list_by_user("test-user-1", page=2, page_size=3)
assert len(items2) == 2
# ── PointsOrderRepository ──
class TestPointsOrderRepository:
def test_create_and_get(self, order_repo):
order = PointsOrder.create(
user_id="test-user-1",
order_type="points",
product_code="starter_pack",
amount_cents=990,
points_amount=100,
)
order_repo.create(order)
fetched = order_repo.get(order.id)
assert fetched is not None
assert fetched.product_code == "starter_pack"
assert fetched.amount_cents == 990
def test_get_nonexistent(self, order_repo):
assert order_repo.get("nonexistent") is None
def test_update_status(self, order_repo):
order = PointsOrder.create(
user_id="test-user-1",
order_type="points",
product_code="starter_pack",
amount_cents=990,
)
order_repo.create(order)
now = datetime.now(timezone.utc)
updated = order_repo.update_status(order.id, "paid", payment_id="pay-123", paid_at=now)
assert updated is not None
assert updated.status == "paid"
assert updated.payment_id == "pay-123"
def test_update_status_nonexistent(self, order_repo):
result = order_repo.update_status("nonexistent", "paid")
assert result is None
def test_list_by_user(self, order_repo):
for ot in ["points", "membership", "points"]:
order = PointsOrder.create(
user_id="test-user-1",
order_type=ot,
product_code="test",
amount_cents=100,
)
order_repo.create(order)
items, total = order_repo.list_by_user("test-user-1", order_type="points")
assert total == 2
def test_list_by_user_with_status(self, order_repo):
order = PointsOrder.create(
user_id="test-user-1",
order_type="points",
product_code="test",
amount_cents=100,
)
order_repo.create(order)
items, total = order_repo.list_by_user("test-user-1", status="pending")
assert total == 1
items2, total2 = order_repo.list_by_user("test-user-1", status="paid")
assert total2 == 0
# ── DailyUsageRepository ──
class TestDailyUsageRepository:
def _today(self):
# The model column is DateTime, so use datetime for comparison
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
return now.replace(hour=0, minute=0, second=0, microsecond=0)
def test_create_and_get(self, daily_repo):
today = self._today()
record = DailyUsageRecord.create(user_id="test-user-1", usage_date=today)
record.count = 1
daily_repo.create(record)
fetched = daily_repo.get_by_user_and_date("test-user-1", today)
assert fetched is not None
assert fetched.count == 1
def test_get_nonexistent(self, daily_repo):
from datetime import timedelta
tomorrow = self._today() + timedelta(days=1)
result = daily_repo.get_by_user_and_date("test-user-1", tomorrow)
assert result is None
def test_update_count(self, daily_repo):
today = self._today()
record = DailyUsageRecord.create(user_id="test-user-1", usage_date=today)
record.count = 1
daily_repo.create(record)
record.count = 3
daily_repo.update_count(record)
fetched = daily_repo.get_by_user_and_date("test-user-1", today)
assert fetched.count == 3
def test_upsert_create(self, daily_repo):
today = self._today()
result = daily_repo.upsert("test-user-1", today, "free_clip")
assert result.count == 1
def test_upsert_increment(self, daily_repo):
today = self._today()
daily_repo.upsert("test-user-1", today, "free_clip")
result = daily_repo.upsert("test-user-1", today, "free_clip")
assert result.count == 2
+140
View File
@@ -0,0 +1,140 @@
"""积分消耗规则单元测试 (#1895)"""
from __future__ import annotations
import math
import pytest
from packages.domain.points_rules import (
DAILY_FREE_CLIP_LIMIT,
FREE_USER_MULTIPLIER,
MEMBER_DISCOUNT,
MEMBERSHIP_PRICES,
POINTS_PACKAGES,
POINTS_SCENES,
calculate_points_cost,
)
class TestPointsScenesConfig:
"""场景配置完整性"""
def test_all_nine_scenes_defined(self):
assert len(POINTS_SCENES) == 9
def test_required_keys_present(self):
for key, scene in POINTS_SCENES.items():
assert "base_points" in scene, f"{key} missing base_points"
assert "unit" in scene, f"{key} missing unit"
assert "name" in scene, f"{key} missing name"
def test_voice_clone_train_is_free(self):
assert POINTS_SCENES["voice_clone_train"]["base_points"] == 0
def test_ai_video_has_extra_per_30s(self):
assert POINTS_SCENES["ai_video"]["extra_per_30s"] == 1
class TestPointsPackages:
def test_three_packages(self):
assert len(POINTS_PACKAGES) == 3
assert POINTS_PACKAGES["starter_pack"]["points"] == 100
assert POINTS_PACKAGES["basic_pack"]["points"] == 500
assert POINTS_PACKAGES["pro_pack"]["points"] == 2000
class TestMembershipPrices:
def test_three_plans(self):
assert len(MEMBERSHIP_PRICES) == 3
assert MEMBERSHIP_PRICES["monthly"]["price_cents"] == 1990
assert MEMBERSHIP_PRICES["quarterly"]["duration_days"] == 90
assert MEMBERSHIP_PRICES["yearly"]["price_cents"] == 15900
class TestDailyFreeLimit:
def test_limit_is_2(self):
assert DAILY_FREE_CLIP_LIMIT == 2
class TestCalculatePointsCost:
"""核心计费逻辑"""
# ── 按次计费 ──
def test_per_time_base_cost(self):
# ai_rewrite: 1积分/次,免费用户 ceil(1 * 1.15) = 2
cost = calculate_points_cost("ai_rewrite", is_member=False, quantity=1)
assert cost == math.ceil(1 * FREE_USER_MULTIPLIER)
def test_per_time_multiple(self):
# ai_cover: 1积分/张,3张 → base=3, free: ceil(3*1.15)=4
cost = calculate_points_cost("ai_cover", is_member=False, quantity=3)
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
# ── 按时长计费 ──
def test_per_minute_base(self):
# ai_voice: 1积分/分钟,3分钟 → base=3, free: ceil(3*1.15)=4
cost = calculate_points_cost("ai_voice", is_member=False, duration_minutes=3)
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
def test_per_minute_rounds_up(self):
# 2.3分钟 → ceil(2.3)=3分钟 → base=3
cost = calculate_points_cost("ai_voice", is_member=False, duration_minutes=2.3)
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
def test_digital_human_expensive(self):
# ai_digital_human: 15积分/分钟,1分钟 → base=15, free: ceil(15*1.15)=18
cost = calculate_points_cost("ai_digital_human", is_member=False, duration_minutes=1)
assert cost == 18
# ── 免费场景 ──
def test_voice_clone_train_free(self):
cost = calculate_points_cost("voice_clone_train", is_member=False)
assert cost == 0
def test_voice_clone_train_free_for_member(self):
cost = calculate_points_cost("voice_clone_train", is_member=True)
assert cost == 0
# ── 混剪额外逻辑 ──
def test_ai_video_short_no_extra(self):
# 20s (0.33min) ≤ 30s,不额外加积分,base=3, free: ceil(3*1.15)=4
cost = calculate_points_cost("ai_video", is_member=False, quantity=1, duration_minutes=0.33)
assert cost == math.ceil(3 * FREE_USER_MULTIPLIER)
def test_ai_video_long_extra_charge(self):
# 80s → base=3 + extra ceil((80-30)/30)=2 → total_base=5, free: ceil(5*1.15)=6
cost = calculate_points_cost("ai_video", is_member=False, quantity=1, duration_minutes=80 / 60)
assert cost == math.ceil(5 * FREE_USER_MULTIPLIER)
# ── 会员折扣 ──
def test_monthly_member_discount(self):
# ai_voice 1分钟 base=1, 月卡0.9 → floor(1*0.9)=1 → max(1,1)=1
cost = calculate_points_cost("ai_voice", is_member=True, duration_minutes=1, member_type="monthly")
assert cost == max(1, math.floor(1 * 0.9))
def test_yearly_member_deep_discount(self):
# ai_digital_human 2分钟 base=30, 年卡0.8 → floor(30*0.8)=24
cost = calculate_points_cost(
"ai_digital_human",
is_member=True,
duration_minutes=2,
member_type="yearly",
)
assert cost == max(1, math.floor(30 * 0.8))
def test_member_without_type_no_discount(self):
# is_member=True 但没传 member_type → 不按会员折扣
cost = calculate_points_cost("ai_voice", is_member=True, duration_minutes=1)
assert cost == 1 # base=1, no discount applied
# ── 异常 ──
def test_unknown_scene_raises(self):
with pytest.raises(ValueError, match="Unknown points scene"):
calculate_points_cost("nonexistent_scene", is_member=False)
+187
View File
@@ -0,0 +1,187 @@
"""PointsService 单元测试 (#1895) — 使用 SQLite 内存数据库"""
from __future__ import annotations
import uuid
from datetime import datetime, timezone
from unittest.mock import patch
import pytest
from sqlalchemy import create_engine, event
from sqlalchemy.orm import Session, sessionmaker
from packages.domain.points_service import PointsService
@pytest.fixture()
def db_session():
"""创建 SQLite 内存数据库 session,包含所有积分相关表。"""
from packages.adapters.sqlalchemy_impl.models import Base
engine = create_engine("sqlite://", echo=False)
# SQLite 不支持 WITH FOR UPDATEmock 掉
@event.listens_for(engine, "connect")
def _disable_for_update(dbapi_conn, connection_record):
pass
Base.metadata.create_all(engine)
SessionLocal = sessionmaker(bind=engine)
session = SessionLocal()
yield session
session.close()
@pytest.fixture()
def service():
return PointsService()
@pytest.fixture()
def user_id():
return uuid.uuid4().hex
class TestGetOrCreateAccount:
def test_creates_new_account(self, service, db_session, user_id):
data = service.get_or_create_account(user_id, db_session)
assert data["user_id"] == user_id
assert data["balance"] == 0
assert data["total_earned"] == 0
assert data["total_spent"] == 0
def test_returns_existing_account(self, service, db_session, user_id):
service.get_or_create_account(user_id, db_session)
data = service.get_or_create_account(user_id, db_session)
assert data["user_id"] == user_id
assert data["balance"] == 0
class TestCheckBalance:
def test_sufficient_when_zero(self, service, db_session, user_id):
result = service.check_balance(user_id, 0, db_session)
assert result["sufficient"] is True
def test_insufficient_when_new_account(self, service, db_session, user_id):
result = service.check_balance(user_id, 10, db_session)
assert result["sufficient"] is False
assert result["remaining_after"] == -10
class TestDeductPoints:
def test_deduct_fails_insufficient_balance(self, service, db_session, user_id):
result = service.deduct_points(user_id, 100, "ai_voice", db_session)
assert result["success"] is False
assert result["transaction_id"] is None
def test_deduct_after_recharge(self, service, db_session, user_id):
# 先充值
service.add_points(user_id, 50, "recharge", db_session)
# 再扣减
result = service.deduct_points(user_id, 20, "ai_voice", db_session)
assert result["success"] is True
assert result["balance"] == 30
def test_deduct_creates_transaction(self, service, db_session, user_id):
service.add_points(user_id, 100, "recharge", db_session)
result = service.deduct_points(user_id, 30, "ai_voice", db_session)
assert result["success"] is True
txns = service.get_transactions(user_id, db_session)
assert txns["total"] == 2 # 1 add + 1 deduct
deduct_txn = [t for t in txns["items"] if t["type"] == "deduct"][0]
assert deduct_txn["amount"] == 30
assert deduct_txn["balance_after"] == 70
class TestAddPoints:
def test_add_new_account(self, service, db_session, user_id):
result = service.add_points(user_id, 100, "recharge:starter_pack", db_session)
assert result["success"] is True
assert result["balance"] == 100
def test_add_accumulates(self, service, db_session, user_id):
service.add_points(user_id, 50, "recharge", db_session)
result = service.add_points(user_id, 30, "bonus", db_session)
assert result["balance"] == 80
class TestRefundPoints:
def test_refund_adds_back(self, service, db_session, user_id):
service.add_points(user_id, 100, "recharge", db_session)
service.deduct_points(user_id, 20, "ai_voice", db_session)
result = service.refund_points(user_id, 20, "ai_voice", db_session)
assert result["success"] is True
assert result["balance"] == 100
def test_refund_creates_refund_transaction(self, service, db_session, user_id):
service.add_points(user_id, 100, "recharge", db_session)
service.refund_points(user_id, 10, "ai_rewrite", db_session)
txns = service.get_transactions(user_id, db_session)
refund_txns = [t for t in txns["items"] if t["type"] == "add" and "refund" in t["source"]]
assert len(refund_txns) == 1
assert "refund:" in refund_txns[0]["source"]
class TestGetTransactions:
def test_empty_for_new_user(self, service, db_session, user_id):
result = service.get_transactions(user_id, db_session)
assert result["total"] == 0
assert result["items"] == []
def test_pagination(self, service, db_session, user_id):
for i in range(5):
service.add_points(user_id, 10, f"batch_{i}", db_session)
result = service.get_transactions(user_id, db_session, page=1, page_size=3)
assert result["total"] == 5
assert len(result["items"]) == 3
result2 = service.get_transactions(user_id, db_session, page=2, page_size=3)
assert len(result2["items"]) == 2
class TestGetDailyUsage:
def test_zero_usage(self, service, db_session, user_id):
with patch("packages.domain.points_service._get_redis_client", return_value=None):
result = service.get_daily_usage(user_id, db_session)
assert result["free_clips_used"] == 0
assert result["free_clips_limit"] == 2
assert result["free_clips_remaining"] == 2
assert "reset_at" in result
def test_after_recording(self, service, db_session, user_id):
with patch("packages.domain.points_service._get_redis_client", return_value=None):
service.record_daily_free_clip(user_id, db_session)
result = service.get_daily_usage(user_id, db_session)
assert result["free_clips_used"] == 1
assert result["free_clips_remaining"] == 1
class TestCreateOrder:
def test_points_order(self, service, db_session, user_id):
result = service.create_order(user_id, "points", "starter_pack", db_session)
assert result["order_type"] == "points"
assert result["product_code"] == "starter_pack"
assert result["amount_cents"] == 990
assert result["status"] == "pending"
def test_membership_order(self, service, db_session, user_id):
result = service.create_order(user_id, "membership", "monthly", db_session)
assert result["order_type"] == "membership"
assert result["amount_cents"] == 1990
def test_unknown_package_raises(self, service, db_session, user_id):
with pytest.raises(ValueError, match="Unknown points package"):
service.create_order(user_id, "points", "nonexistent", db_session)
def test_unknown_membership_raises(self, service, db_session, user_id):
with pytest.raises(ValueError, match="Unknown membership type"):
service.create_order(user_id, "membership", "lifetime", db_session)
def test_unknown_order_type_raises(self, service, db_session, user_id):
with pytest.raises(ValueError, match="Unknown order type"):
service.create_order(user_id, "insurance", "basic", db_session)