feat: 会员+积分系统后端 (#1895) #1919
@@ -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")
|
||||
@@ -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"],
|
||||
)
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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_points,source 前缀 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 记录每日额度")
|
||||
|
||||
# 降级/兜底到 DB(upsert 语义)
|
||||
_, _, _, 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": "确认支付异常"}
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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: ...
|
||||
@@ -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]: ...
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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 UPDATE,mock 掉
|
||||
@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)
|
||||
Reference in New Issue
Block a user