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