diff --git a/alembic/versions/076_membership_points.py b/alembic/versions/076_membership_points.py new file mode 100644 index 000000000..de3e5b95f --- /dev/null +++ b/alembic/versions/076_membership_points.py @@ -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") diff --git a/apps/api/app/api/router.py b/apps/api/app/api/router.py index 3dd8f9afc..fbe62c04c 100755 --- a/apps/api/app/api/router.py +++ b/apps/api/app/api/router.py @@ -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"], +) diff --git a/apps/api/app/api/routes/points.py b/apps/api/app/api/routes/points.py new file mode 100644 index 000000000..9bdfb48d3 --- /dev/null +++ b/apps/api/app/api/routes/points.py @@ -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 diff --git a/apps/api/app/schemas/points.py b/apps/api/app/schemas/points.py new file mode 100644 index 000000000..0ab25b8df --- /dev/null +++ b/apps/api/app/schemas/points.py @@ -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 diff --git a/migrations/008_membership_points.sql b/migrations/008_membership_points.sql new file mode 100644 index 000000000..2d9d35001 --- /dev/null +++ b/migrations/008_membership_points.sql @@ -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); diff --git a/packages/adapters/sqlalchemy_impl/daily_usage_repository.py b/packages/adapters/sqlalchemy_impl/daily_usage_repository.py new file mode 100644 index 000000000..d4996b4fd --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/daily_usage_repository.py @@ -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, + ) diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 270f0891e..648aa9ca8 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -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)) diff --git a/packages/adapters/sqlalchemy_impl/points_account_repository.py b/packages/adapters/sqlalchemy_impl/points_account_repository.py new file mode 100644 index 000000000..29c391caa --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/points_account_repository.py @@ -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, + ) diff --git a/packages/adapters/sqlalchemy_impl/points_order_repository.py b/packages/adapters/sqlalchemy_impl/points_order_repository.py new file mode 100644 index 000000000..59d0ce244 --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/points_order_repository.py @@ -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, + ) diff --git a/packages/adapters/sqlalchemy_impl/points_transaction_repository.py b/packages/adapters/sqlalchemy_impl/points_transaction_repository.py new file mode 100644 index 000000000..e98483c6d --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/points_transaction_repository.py @@ -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, + ) diff --git a/packages/domain/__init__.py b/packages/domain/__init__.py index 39552d894..6ed2f89ce 100755 --- a/packages/domain/__init__.py +++ b/packages/domain/__init__.py @@ -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", diff --git a/packages/domain/daily_usage_record.py b/packages/domain/daily_usage_record.py new file mode 100644 index 000000000..ba4ccb197 --- /dev/null +++ b/packages/domain/daily_usage_record.py @@ -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, + ) diff --git a/packages/domain/points_account.py b/packages/domain/points_account.py new file mode 100644 index 000000000..d6374c328 --- /dev/null +++ b/packages/domain/points_account.py @@ -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) diff --git a/packages/domain/points_order.py b/packages/domain/points_order.py new file mode 100644 index 000000000..0666179c8 --- /dev/null +++ b/packages/domain/points_order.py @@ -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, + ) diff --git a/packages/domain/points_rules.py b/packages/domain/points_rules.py new file mode 100644 index 000000000..25aff4c1d --- /dev/null +++ b/packages/domain/points_rules.py @@ -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 diff --git a/packages/domain/points_service.py b/packages/domain/points_service.py new file mode 100644 index 000000000..d6caba734 --- /dev/null +++ b/packages/domain/points_service.py @@ -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": "确认支付异常"} diff --git a/packages/domain/points_transaction.py b/packages/domain/points_transaction.py new file mode 100644 index 000000000..93defbbc6 --- /dev/null +++ b/packages/domain/points_transaction.py @@ -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, + ) diff --git a/packages/middleware/__init__.py b/packages/middleware/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/packages/middleware/points_gate.py b/packages/middleware/points_gate.py new file mode 100644 index 000000000..4d4770cdb --- /dev/null +++ b/packages/middleware/points_gate.py @@ -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) diff --git a/packages/ports/daily_usage_repository.py b/packages/ports/daily_usage_repository.py new file mode 100644 index 000000000..e7fa12fe4 --- /dev/null +++ b/packages/ports/daily_usage_repository.py @@ -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: ... diff --git a/packages/ports/points_account_repository.py b/packages/ports/points_account_repository.py new file mode 100644 index 000000000..e2c12ba12 --- /dev/null +++ b/packages/ports/points_account_repository.py @@ -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: ... diff --git a/packages/ports/points_order_repository.py b/packages/ports/points_order_repository.py new file mode 100644 index 000000000..daff5805d --- /dev/null +++ b/packages/ports/points_order_repository.py @@ -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]: ... diff --git a/packages/ports/points_transaction_repository.py b/packages/ports/points_transaction_repository.py new file mode 100644 index 000000000..1c4f62081 --- /dev/null +++ b/packages/ports/points_transaction_repository.py @@ -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]: ... diff --git a/tests/unit/test_points_gate.py b/tests/unit/test_points_gate.py new file mode 100644 index 000000000..44c8914be --- /dev/null +++ b/tests/unit/test_points_gate.py @@ -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 diff --git a/tests/unit/test_points_repositories.py b/tests/unit/test_points_repositories.py new file mode 100644 index 000000000..164664963 --- /dev/null +++ b/tests/unit/test_points_repositories.py @@ -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 diff --git a/tests/unit/test_points_rules.py b/tests/unit/test_points_rules.py new file mode 100644 index 000000000..36e8b9d89 --- /dev/null +++ b/tests/unit/test_points_rules.py @@ -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) diff --git a/tests/unit/test_points_service.py b/tests/unit/test_points_service.py new file mode 100644 index 000000000..33b6816be --- /dev/null +++ b/tests/unit/test_points_service.py @@ -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)