diff --git a/.env.example b/.env.example index 8453460b..affa5e38 100644 --- a/.env.example +++ b/.env.example @@ -40,8 +40,10 @@ AI_IMAGE_MODEL=gemini-2.5-flash-image AI_VIDEO_MODEL=kling-v2-5-turbo # ── 积分定价 ── -QUOTA_REGISTER_GIFT_AMOUNT=100 -QUOTA_INVITE_REWARD_AMOUNT=50 +QUOTA_REGISTER_GIFT_AMOUNT=300 +QUOTA_INVITE_REWARD_AMOUNT=200 +QUOTA_INVITE_REWARD_DAILY_LIMIT=3 +QUOTA_INVITE_CODE_TTL_DAYS=30 QUOTA_GENERATE_IMAGE_COST=10 QUOTA_GENERATE_ACTION_COST=50 diff --git a/backend/packages/app/src/windup_app/bootstrap/app.py b/backend/packages/app/src/windup_app/bootstrap/app.py index 586c8b2d..b08dec5c 100644 --- a/backend/packages/app/src/windup_app/bootstrap/app.py +++ b/backend/packages/app/src/windup_app/bootstrap/app.py @@ -21,8 +21,7 @@ from windup_app.server.character.service import service as character_service from windup_app.server.orchestrator.dispatcher import GenerationDispatcher from windup_app.server.project.model import Project # noqa: F401 -from windup_app.server.quota.model import CreditAccount, CreditTransaction # noqa: F401 -# InviteCode, InviteRecord, TokenUsage 暂不实现 +from windup_app.server.quota.model import CreditAccount, CreditTransaction, InviteCode, InviteRecord # noqa: F401 from windup_app.server.user.model import User # noqa: F401 from windup_app.server.workflow_run.model import WorkflowRun # noqa: F401 from windup_app.web.api.auth import router as auth_router diff --git a/backend/packages/app/src/windup_app/server/quota/interface.py b/backend/packages/app/src/windup_app/server/quota/interface.py index c419c466..25b88e66 100644 --- a/backend/packages/app/src/windup_app/server/quota/interface.py +++ b/backend/packages/app/src/windup_app/server/quota/interface.py @@ -10,6 +10,8 @@ from windup_app.server.quota.model import ( CreditAccountView, CreditTransactionView, + InviteCode, + InviteCodeView, ) @@ -35,7 +37,12 @@ def reserve_credit( @abstractmethod def capture_credit( - self, session: Session, user_id: int, actual_amount: int, ref_id: str, frozen_amount: int + self, + session: Session, + user_id: int, + actual_amount: int, + ref_id: str, + frozen_amount: int, ) -> None: """预付费扣减:冻结转消耗。 @@ -65,7 +72,12 @@ def release_credit( @abstractmethod def credit( - self, session: Session, user_id: int, amount: int, reason: int, ref_id: str | None = None + self, + session: Session, + user_id: int, + amount: int, + reason: int, + ref_id: str | None = None, ) -> None: """入账:增加可用余额与累计获得。""" @@ -78,18 +90,22 @@ def list_transactions( """分页查询积分流水,返回 (列表, 总数)。""" # -- 邀请码 ----------------------------------------------------------- - # TODO 目前先不实现。 - # @abstractmethod - # def get_invite_code(self, session: Session, user_id: int) -> InviteCodeView | None: - # """获取用户当前邀请码。""" - # - # @abstractmethod - # def generate_invite_code(self, session: Session, user_id: int) -> InviteCodeView: - # """生成新邀请码(替换旧码)。""" - # - # @abstractmethod - # def redeem_invite_code(self, session: Session, user_id: int, code: str) -> None: - # """兑换邀请码,双方各得积分。 - # - # :raises BizException: 邀请码无效 / 已达上限 / 已填过码。 - # """ + + @abstractmethod + def get_invite_code(self, session: Session, user_id: int) -> InviteCodeView: + """获取当前未过期邀请码;没有或已过期则签发新行。""" + + @abstractmethod + def generate_invite_code(self, session: Session, user_id: int) -> InviteCodeView: + """签发新邀请码:插入新行,仍有效的旧码立即过期但保留。""" + + @abstractmethod + def require_active_invite(self, session: Session, code: str) -> InviteCode: + """注册前校验邀请码存在且未过期。非法返回「邀请码无效」,过期返回「邀请码已过期」。""" + + @abstractmethod + def redeem_invite_code(self, session: Session, user_id: int, code: str) -> None: + """注册时兑换邀请码。被邀请人始终得邀请奖励;邀请人受每日人数上限。 + + :raises BizException: 邀请码无效 / 已过期 / 已填过码 / 不能填自己的码。 + """ diff --git a/backend/packages/app/src/windup_app/server/quota/model.py b/backend/packages/app/src/windup_app/server/quota/model.py index 45e82ff4..e04a0ff3 100644 --- a/backend/packages/app/src/windup_app/server/quota/model.py +++ b/backend/packages/app/src/windup_app/server/quota/model.py @@ -15,11 +15,19 @@ """ from dataclasses import dataclass, field -from datetime import datetime, timezone - -from sqlalchemy import BigInteger, DateTime, Integer, SmallInteger, String, UniqueConstraint +from datetime import datetime, timedelta, timezone + +from sqlalchemy import ( + BigInteger, + DateTime, + Integer, + SmallInteger, + String, + UniqueConstraint, +) from sqlalchemy.orm import Mapped, mapped_column +from windup_framework.config.quota import settings as quota_settings from windup_framework.db import Base @@ -107,18 +115,70 @@ class CreditTransaction(Base): ) -# -- 以下 ORM 暂不实现(枚举 / 接口已预留)---------------------------------- -# -# class InviteCode(Base): -# """邀请码。""" -# __tablename__ = "windup_invite_code" -# ... -# -# class InviteRecord(Base): -# """邀请记录。""" -# __tablename__ = "windup_invite_record" -# ... -# +class InviteCode(Base): + """用户邀请码。只增不删;轮换插入新行,旧行保留。""" + + __tablename__ = "windup_invite_code" + + id: Mapped[int] = mapped_column( + BigInteger().with_variant(Integer, "sqlite"), + primary_key=True, + autoincrement=True, + ) + user_id: Mapped[int] = mapped_column( + BigInteger().with_variant(Integer, "sqlite"), + index=True, + nullable=False, + ) + code: Mapped[str] = mapped_column(String(16), unique=True, nullable=False) + used_count: Mapped[int] = mapped_column(Integer, nullable=False, default=0) + expires_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + nullable=False, + default=lambda: datetime.now(timezone.utc) + + timedelta(days=quota_settings.invite_code_ttl_days), + ) + create_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + nullable=False, + default=lambda: datetime.now(timezone.utc), + ) + update_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + nullable=False, + default=lambda: datetime.now(timezone.utc), + onupdate=lambda: datetime.now(timezone.utc), + ) + + +class InviteRecord(Base): + """一次成功的邀请关系。被邀请人只能出现一次。""" + + __tablename__ = "windup_invite_record" + + id: Mapped[int] = mapped_column( + BigInteger().with_variant(Integer, "sqlite"), + primary_key=True, + autoincrement=True, + ) + inviter_id: Mapped[int] = mapped_column( + BigInteger().with_variant(Integer, "sqlite"), + nullable=False, + index=True, + ) + invitee_id: Mapped[int] = mapped_column( + BigInteger().with_variant(Integer, "sqlite"), + unique=True, + nullable=False, + ) + code: Mapped[str] = mapped_column(String(16), nullable=False) + create_at: Mapped[datetime] = mapped_column( + DateTime(timezone=True), + nullable=False, + default=lambda: datetime.now(timezone.utc), + ) + + # class TokenUsage(Base): # """Token 用量记录。""" # __tablename__ = "windup_token_usage" @@ -156,9 +216,12 @@ class CreditTransactionView: create_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) -# -- 暂不实现 -- -# -# @dataclass -# class InviteCodeView: -# """邀请码视图。""" -# ... +@dataclass +class InviteCodeView: + """邀请码视图。""" + + code: str + used_count: int = 0 + expires_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) + create_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) + update_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) diff --git a/backend/packages/app/src/windup_app/server/quota/service.py b/backend/packages/app/src/windup_app/server/quota/service.py index 5fcb63bf..e7a8d72f 100644 --- a/backend/packages/app/src/windup_app/server/quota/service.py +++ b/backend/packages/app/src/windup_app/server/quota/service.py @@ -12,8 +12,12 @@ """ import logging +import re +import secrets +from datetime import datetime, timedelta, timezone from sqlalchemy import func, select +from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session from windup_common.enums.biz_code import BizCode @@ -26,10 +30,68 @@ CreditAccountView, CreditTransaction, CreditTransactionView, + InviteCode, + InviteCodeView, + InviteRecord, ) +from windup_app.server.user.model import User +from windup_framework.config.quota import settings as quota_settings logger = logging.getLogger("windup.quota.service") +_INVITE_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" +_INVITE_CODE_LENGTH = 8 +_INVITE_CODE_RE = re.compile(rf"^[{re.escape(_INVITE_ALPHABET)}]{{4,16}}$") + + +def normalize_invite_code(code: str) -> str: + return code.strip().upper() + + +def parse_invite_code(code: str) -> str: + """解析邀请链接传入的邀请码,字符集与前端 INVITE_CODE_PATTERN 一致。""" + normalized = normalize_invite_code(code) + if _INVITE_CODE_RE.fullmatch(normalized) is None: + raise BizException("邀请码无效", code=BizCode.BAD_REQUEST) + return normalized + + +def _new_invite_code() -> str: + return "".join(secrets.choice(_INVITE_ALPHABET) for _ in range(_INVITE_CODE_LENGTH)) + + +def _is_invitee_unique_violation(exc: IntegrityError) -> bool: + text = f"{getattr(exc, 'orig', '')} {exc}".lower() + return "invitee" in text or "windup_invite_record" in text + + +def _now() -> datetime: + return datetime.now(timezone.utc) + + +def _utc_day_start(now: datetime | None = None) -> datetime: + current = now or _now() + if current.tzinfo is None: + current = current.replace(tzinfo=timezone.utc) + return current.astimezone(timezone.utc).replace( + hour=0, minute=0, second=0, microsecond=0 + ) + + +def _is_expired(expires_at: datetime) -> bool: + exp = expires_at if expires_at.tzinfo else expires_at.replace(tzinfo=timezone.utc) + return exp <= _now() + + +def _to_invite_view(row: InviteCode) -> InviteCodeView: + return InviteCodeView( + code=row.code, + used_count=row.used_count, + expires_at=row.expires_at, + create_at=row.create_at, + update_at=row.update_at, + ) + def _to_account_view(account: CreditAccount) -> CreditAccountView: return CreditAccountView( @@ -120,17 +182,30 @@ def reserve_credit( session.flush() self._write_txn( - session, user_id, -amount, CreditReason.FROZEN, - BillingMode.PREPAID, account.balance, ref_id, + session, + user_id, + -amount, + CreditReason.FROZEN, + BillingMode.PREPAID, + account.balance, + ref_id, ) logger.info( "[WINDUP] 积分冻结 | user_id=%s amount=%s ref_id=%s balance=%s", - user_id, amount, ref_id, account.balance, + user_id, + amount, + ref_id, + account.balance, ) def capture_credit( - self, session: Session, user_id: int, actual_amount: int, ref_id: str, frozen_amount: int + self, + session: Session, + user_id: int, + actual_amount: int, + ref_id: str, + frozen_amount: int, ) -> None: """预付费扣减:frozen -= frozen_amount, total_spent += actual_amount。 @@ -158,20 +233,34 @@ def capture_credit( # 写扣减流水 self._write_txn( - session, user_id, -actual_amount, CreditReason.CAPTURED, - BillingMode.PREPAID, account.balance, ref_id, + session, + user_id, + -actual_amount, + CreditReason.CAPTURED, + BillingMode.PREPAID, + account.balance, + ref_id, ) # 有差额退回时写退款流水(用不同 reason 区分,ref_id 加后缀去重) if refund > 0: self._write_txn( - session, user_id, refund, CreditReason.REFUND, - BillingMode.PREPAID, account.balance, f"{ref_id}:refund", + session, + user_id, + refund, + CreditReason.REFUND, + BillingMode.PREPAID, + account.balance, + f"{ref_id}:refund", ) logger.info( "[WINDUP] 积分扣减 | user_id=%s actual=%s frozen=%s refund=%s balance=%s", - user_id, actual_amount, frozen_amount, refund, account.balance, + user_id, + actual_amount, + frozen_amount, + refund, + account.balance, ) def release_credit( @@ -191,13 +280,21 @@ def release_credit( session.flush() self._write_txn( - session, user_id, amount, CreditReason.REFUND, - BillingMode.PREPAID, account.balance, f"{ref_id}:release", + session, + user_id, + amount, + CreditReason.REFUND, + BillingMode.PREPAID, + account.balance, + f"{ref_id}:release", ) logger.info( "[WINDUP] 积分解冻 | user_id=%s amount=%s ref_id=%s balance=%s", - user_id, amount, ref_id, account.balance, + user_id, + amount, + ref_id, + account.balance, ) # -- 后付费:原子扣减(暂不实现,AGENT_TOKEN / POSTPAID 枚举已预留)------ @@ -211,7 +308,12 @@ def release_credit( # -- 入账(赠送 / 奖励 / 管理员调整)---------------------------------- def credit( - self, session: Session, user_id: int, amount: int, reason: int, ref_id: str | None = None + self, + session: Session, + user_id: int, + amount: int, + reason: int, + ref_id: str | None = None, ) -> None: """入账:balance += amount, total_earned += amount。""" if amount <= 0: @@ -223,13 +325,21 @@ def credit( session.flush() self._write_txn( - session, user_id, amount, reason, - BillingMode.PREPAID, account.balance, ref_id, + session, + user_id, + amount, + reason, + BillingMode.PREPAID, + account.balance, + ref_id, ) logger.info( "[WINDUP] 积分入账 | user_id=%s amount=%s reason=%s balance=%s", - user_id, amount, reason, account.balance, + user_id, + amount, + reason, + account.balance, ) # -- 流水查询 --------------------------------------------------------- @@ -254,8 +364,128 @@ def list_transactions( return [_to_txn_view(r) for r in rows], total or 0 - # -- 邀请码(暂不实现)------------------------------------------------- - # TODO: get_invite_code / generate_invite_code / redeem_invite_code + # -- 邀请码 ----------------------------------------------------------- + + def get_invite_code(self, session: Session, user_id: int) -> InviteCodeView: + row = session.scalar( + select(InviteCode) + .where(InviteCode.user_id == user_id, InviteCode.expires_at > _now()) + .order_by(InviteCode.id.desc()) + ) + if row is not None: + return _to_invite_view(row) + return self.generate_invite_code(session, user_id) + + def generate_invite_code(self, session: Session, user_id: int) -> InviteCodeView: + if session.get(User, user_id) is None: + raise BizException("用户不存在", code=BizCode.NOT_FOUND) + + now = _now() + existing = session.scalars( + select(InviteCode) + .where(InviteCode.user_id == user_id) + .with_for_update() + ).all() + for row in existing: + if not _is_expired(row.expires_at): + row.expires_at = now + + row = InviteCode( + user_id=user_id, + code=self._allocate_invite_code(session), + used_count=0, + expires_at=now + + timedelta(days=quota_settings.invite_code_ttl_days), + ) + session.add(row) + session.flush() + logger.info("[WINDUP] 生成邀请码 | user_id=%s code=%s", user_id, row.code) + return _to_invite_view(row) + + def _allocate_invite_code(self, session: Session) -> str: + for _ in range(16): + code = _new_invite_code() + if session.scalar(select(InviteCode.id).where(InviteCode.code == code)) is None: + return code + raise BizException("邀请码生成失败,请稍后重试", code=BizCode.BAD_REQUEST) + + def require_active_invite(self, session: Session, code: str) -> InviteCode: + normalized = parse_invite_code(code) + invite = session.scalar( + select(InviteCode).where(InviteCode.code == normalized) + ) + if invite is None: + raise BizException("邀请码无效", code=BizCode.BAD_REQUEST) + if _is_expired(invite.expires_at): + raise BizException("邀请码已过期", code=BizCode.NOT_FOUND) + return invite + + def redeem_invite_code(self, session: Session, user_id: int, code: str) -> None: + invite = self.require_active_invite(session, code) + if invite.user_id == user_id: + raise BizException("不能填写自己的邀请码", code=BizCode.BAD_REQUEST) + if session.get(User, user_id) is None: + raise BizException("用户不存在", code=BizCode.NOT_FOUND) + + existing = session.scalar( + select(InviteRecord.id).where(InviteRecord.invitee_id == user_id) + ) + if existing is not None: + raise BizException("已填写过邀请码", code=BizCode.BAD_REQUEST) + + self._get_account_for_update(session, invite.user_id) + + record = InviteRecord( + inviter_id=invite.user_id, + invitee_id=user_id, + code=invite.code, + ) + session.add(record) + invite.used_count += 1 + try: + session.flush() + except IntegrityError as exc: + if _is_invitee_unique_violation(exc): + raise BizException("已填写过邀请码", code=BizCode.BAD_REQUEST) from exc + raise + + reward = quota_settings.invite_reward_amount + today_count = session.scalar( + select(func.count()) + .select_from(InviteRecord) + .where( + InviteRecord.inviter_id == invite.user_id, + InviteRecord.create_at >= _utc_day_start(), + ) + ) or 0 + if today_count <= quota_settings.invite_reward_daily_limit: + self.credit( + session, + invite.user_id, + reward, + int(CreditReason.INVITE_REWARD), + f"invite:{user_id}:inviter", + ) + else: + logger.info( + "[WINDUP] 邀请人日限额已满,跳过邀请人奖励 | inviter=%s invitee=%s count=%s", + invite.user_id, + user_id, + today_count, + ) + self.credit( + session, + user_id, + reward, + int(CreditReason.INVITE_REWARD), + f"invite:{user_id}:invitee", + ) + logger.info( + "[WINDUP] 兑换邀请码 | invitee=%s inviter=%s code=%s", + user_id, + invite.user_id, + invite.code, + ) service = SqlAlchemyQuotaService() diff --git a/backend/packages/app/src/windup_app/server/user/interface.py b/backend/packages/app/src/windup_app/server/user/interface.py index c8c0d0d2..6ea3d826 100644 --- a/backend/packages/app/src/windup_app/server/user/interface.py +++ b/backend/packages/app/src/windup_app/server/user/interface.py @@ -27,9 +27,9 @@ class UserService(ABC): @abstractmethod def register_by_email(self, session: Session, input: RegisterInput) -> LoginResult: - """邮箱+密码注册,注册成功即登录。 + """邮箱+验证码+密码注册。邀请码选填。 - :raises windup_common.exceptions.BizException: 邮箱已注册。 + :raises windup_common.exceptions.BizException: 邮箱已注册 / 邀请码无效。 """ # -- 登录 ------------------------------------------------------------ @@ -53,9 +53,9 @@ def send_verification_code(self, email: str, purpose: str) -> None: @abstractmethod def login_by_code(self, session: Session, input: LoginByCodeInput) -> LoginResult: - """邮箱+验证码登录。内测期间不自动建号。 + """邮箱+验证码登录。未知邮箱自动建号并赠送注册积分。 - :raises windup_common.exceptions.BizException: 验证码错误 / 已过期 / 账号不存在 / 账号已封禁。 + :raises windup_common.exceptions.BizException: 验证码错误 / 已过期 / 账号已封禁。 """ # -- 登出 ------------------------------------------------------------ diff --git a/backend/packages/app/src/windup_app/server/user/model.py b/backend/packages/app/src/windup_app/server/user/model.py index 31b51464..c7460228 100644 --- a/backend/packages/app/src/windup_app/server/user/model.py +++ b/backend/packages/app/src/windup_app/server/user/model.py @@ -115,6 +115,7 @@ class RegisterInput: password: str code: str nickname: str | None = None + invite_code: str | None = None @dataclass diff --git a/backend/packages/app/src/windup_app/server/user/service.py b/backend/packages/app/src/windup_app/server/user/service.py index d5bc1b03..572cf099 100644 --- a/backend/packages/app/src/windup_app/server/user/service.py +++ b/backend/packages/app/src/windup_app/server/user/service.py @@ -49,7 +49,7 @@ JWT_SECRET = jwt_settings.secret.get_secret_value() JWT_ALGORITHM = "HS256" -ACCESS_TOKEN_EXPIRE_SECONDS = 15 * 60 # 15 分钟 +ACCESS_TOKEN_EXPIRE_SECONDS = 15 * 60 # 15 分钟 REFRESH_TOKEN_EXPIRE_SECONDS = 7 * 24 * 3600 # 7 天 # -- 密码哈希 ------------------------------------------------------------- @@ -64,6 +64,7 @@ def _verify_password(password: str, hashed: str) -> bool: """验证密码。""" return bcrypt.checkpw(password.encode(), hashed.encode()) + # -- Redis key 前缀 ------------------------------------------------------- VERIFY_COOLDOWN_KEY = "verify:cooldown:{email}" @@ -72,12 +73,12 @@ def _verify_password(password: str, hashed: str) -> bool: LOGIN_FAIL_KEY = "login:fail:{email}" LOGIN_LOCK_KEY = "login:lock:{email}" -VERIFY_CODE_TTL = 300 # 5 分钟 -COOLDOWN_TTL = 60 # 60 秒 +VERIFY_CODE_TTL = 300 # 5 分钟 +COOLDOWN_TTL = 60 # 60 秒 -LOGIN_FAIL_LIMIT = 5 # 连续错误密码上限 -LOGIN_FAIL_WINDOW = 15 * 60 # 失败计数窗口 15 分钟 -LOGIN_LOCK_DURATION = 15 * 60 # 锁定时长 15 分钟 +LOGIN_FAIL_LIMIT = 5 # 连续错误密码上限 +LOGIN_FAIL_WINDOW = 15 * 60 # 失败计数窗口 15 分钟 +LOGIN_LOCK_DURATION = 15 * 60 # 锁定时长 15 分钟 def _hash_token(token: str) -> str: @@ -166,11 +167,18 @@ def redis(self) -> redis_lib.Redis: # -- 注册 ------------------------------------------------------------ - def register_by_email( - self, session: Session, input: RegisterInput - ) -> LoginResult: - """邮箱+验证码+密码注册。""" - # 校验验证码 + def register_by_email(self, session: Session, input: RegisterInput) -> LoginResult: + """邮箱+验证码+密码注册。邀请码选填。""" + from windup_app.server.quota.service import ( + parse_invite_code, + service as quota_service, + ) + + raw_invite = (input.invite_code or "").strip() + invite_code = parse_invite_code(raw_invite) if raw_invite else None + if invite_code is not None: + quota_service.require_active_invite(session, invite_code) + self._verify_code(input.email, input.code, "register") # 检查邮箱唯一 @@ -184,13 +192,17 @@ def register_by_email( email=input.email, password_hash=_hash_password(input.password), nickname=input.nickname, - email_verified_at=datetime.now(timezone.utc), # 注册即验证(已通过验证码校验) + email_verified_at=datetime.now( + timezone.utc + ), # 注册即验证(已通过验证码校验) ) session.add(user) session.flush() - # 注册送积分 + # 注册送积分;有邀请码再发双方邀请奖励 self._create_credit_account(session, user.id) + if invite_code is not None: + quota_service.redeem_invite_code(session, user.id, invite_code) # 注册即登录,签发 token access_token = create_access_token(user.id, user.email) @@ -270,13 +282,12 @@ def login_by_password( def send_verification_code(self, email: str, purpose: str) -> None: """发送邮箱验证码。""" - if purpose == "register": - raise BizException("内测期间暂不开放注册", code=BizCode.BAD_REQUEST) - # 频率限制 cooldown_key = VERIFY_COOLDOWN_KEY.format(email=email) if self.redis.get(cooldown_key): - raise BizException("发送过于频繁,请稍后再试", code=BizCode.TOO_MANY_REQUESTS) + raise BizException( + "发送过于频繁,请稍后再试", code=BizCode.TOO_MANY_REQUESTS + ) code = _generate_code() code_key = VERIFY_CODE_KEY.format(purpose=purpose, email=email) @@ -302,20 +313,26 @@ def _verify_code(self, email: str, code: str, purpose: str) -> None: # 验证通过,删除验证码 self.redis.delete(code_key) - def login_by_code( - self, session: Session, input: LoginByCodeInput - ) -> LoginResult: - """邮箱+验证码登录。内测期间不自动建号。""" + def login_by_code(self, session: Session, input: LoginByCodeInput) -> LoginResult: + """邮箱+验证码登录。未知邮箱自动建号并赠送注册积分。""" # 校验验证码 self._verify_code(input.email, input.code, "login") user = session.scalar(select(User).where(User.email == input.email)) if user is None: - raise BizException("账号不存在", code=BizCode.NOT_FOUND) - if user.status == UserStatus.BANNED: - raise BizException("账号已被封禁", code=BizCode.BAD_REQUEST) - if user.email_verified_at is None: - user.email_verified_at = datetime.now(timezone.utc) + user = User( + email=input.email, + password_hash="", + email_verified_at=datetime.now(timezone.utc), + ) + session.add(user) + session.flush() + self._create_credit_account(session, user.id) + else: + if user.status == UserStatus.BANNED: + raise BizException("账号已被封禁", code=BizCode.BAD_REQUEST) + if user.email_verified_at is None: + user.email_verified_at = datetime.now(timezone.utc) user.last_login_at = datetime.now(timezone.utc) session.flush() @@ -443,9 +460,7 @@ def change_password( self._revoke_all_user_tokens(user_id) logger.info("[WINDUP] 密码已修改 | user_id=%s", user_id) - def reset_password( - self, session: Session, input: ResetPasswordInput - ) -> None: + def reset_password(self, session: Session, input: ResetPasswordInput) -> None: """邮箱+验证码重置密码(忘记密码场景)。""" # 校验验证码(purpose 必须为 reset_password) self._verify_code(input.email, input.code, "reset_password") @@ -517,7 +532,8 @@ def _create_credit_account(self, session: Session, user_id: int) -> None: logger.info( "[WINDUP] 注册送积分 | user_id=%s amount=%s", - user_id, quota_settings.register_gift_amount, + user_id, + quota_settings.register_gift_amount, ) def _store_refresh_token(self, jti: str, user_id: int) -> None: diff --git a/backend/packages/app/src/windup_app/web/api/auth.py b/backend/packages/app/src/windup_app/web/api/auth.py index 72504c2f..eb76d7db 100644 --- a/backend/packages/app/src/windup_app/web/api/auth.py +++ b/backend/packages/app/src/windup_app/web/api/auth.py @@ -6,14 +6,20 @@ import logging from fastapi import APIRouter, Depends, Request -from pydantic import BaseModel, ConfigDict, Field, EmailStr +from pydantic import BaseModel, ConfigDict, Field, EmailStr, field_validator from sqlalchemy.orm import Session from windup_common.result import Response from windup_framework.db import get_session -from windup_app.server.user.model import ResetPasswordInput, UpdateNicknameInput, User, UserView +from windup_app.server.user.model import ( + RegisterInput, + ResetPasswordInput, + UpdateNicknameInput, + User, + UserView, +) from windup_app.server.user.service import service logger = logging.getLogger("windup.auth.api") @@ -31,6 +37,18 @@ class RegisterRequest(BaseModel): password: str = Field(min_length=8, max_length=128) code: str = Field(min_length=6, max_length=6, description="邮箱验证码") nickname: str | None = Field(default=None, max_length=50) + invite_code: str | None = Field( + default=None, + max_length=16, + description="邀请链接中的邀请码,选填;有则发双方邀请奖励", + ) + + @field_validator("invite_code", mode="before") + @classmethod + def blank_invite_code(cls, value: object) -> object: + if isinstance(value, str) and not value.strip(): + return None + return value class LoginRequest(BaseModel): @@ -77,7 +95,9 @@ class ResetPasswordRequest(BaseModel): """重置密码请求(忘记密码场景)。""" email: EmailStr - code: str = Field(min_length=6, max_length=6, description="reset_password 用途的验证码") + code: str = Field( + min_length=6, max_length=6, description="reset_password 用途的验证码" + ) new_password: str = Field(min_length=8, max_length=128) @@ -111,15 +131,25 @@ class UserOut(BaseModel): @router.post("/register", response_model=Response[TokenResponse]) def register(body: RegisterRequest, session: Session = Depends(get_session)): - """邮箱+验证码+密码注册。 - - 内测期间关闭公开注册,路由与请求模型保留以便以后重新开放。 - """ - from windup_common.enums.biz_code import BizCode - from windup_common.exceptions import BizException - - del body, session - raise BizException("内测期间暂不开放注册", code=BizCode.BAD_REQUEST) + """邮箱+验证码+密码注册。邀请码选填。""" + result = service.register_by_email( + session, + RegisterInput( + email=body.email, + password=body.password, + code=body.code, + nickname=body.nickname, + invite_code=body.invite_code, + ), + ) + return Response.success( + TokenResponse( + access_token=result.access_token, + refresh_token=result.refresh_token, + user=result.user, + ), + message="注册成功", + ) @router.post("/login", response_model=Response[TokenResponse]) @@ -127,7 +157,9 @@ def login(body: LoginRequest, session: Session = Depends(get_session)): """邮箱+密码+验证码登录。""" result = service.login_by_password( session, - type("LoginByPasswordInput", (), {"email": body.email, "password": body.password})(), + type( + "LoginByPasswordInput", (), {"email": body.email, "password": body.password} + )(), ) return Response.success( TokenResponse( @@ -148,7 +180,7 @@ def send_code(body: SendCodeRequest): @router.post("/login-by-code", response_model=Response[TokenResponse]) def login_by_code(body: LoginByCodeRequest, session: Session = Depends(get_session)): - """验证码登录。内测期间不自动注册。""" + """验证码登录。未知邮箱自动建号并赠送注册积分。""" result = service.login_by_code( session, type("LoginByCodeInput", (), {"email": body.email, "code": body.code})(), @@ -191,26 +223,37 @@ def get_me(request: Request, session: Session = Depends(get_session)): if user is None: from windup_common.enums.biz_code import BizCode from windup_common.exceptions import BizException + raise BizException("用户不存在", code=BizCode.NOT_FOUND) return Response.success( UserOut( id=user.id, email=user.email, nickname=user.nickname, - email_verified_at=user.email_verified_at.isoformat() if user.email_verified_at else None, + email_verified_at=user.email_verified_at.isoformat() + if user.email_verified_at + else None, status=user.status, ) ) @router.post("/change-password", response_model=Response[None]) -def change_password(body: ChangePasswordRequest, request: Request, session: Session = Depends(get_session)): +def change_password( + body: ChangePasswordRequest, + request: Request, + session: Session = Depends(get_session), +): """修改密码。""" current_user = request.state.current_user service.change_password( session, current_user.id, - type("ChangePasswordInput", (), {"old_password": body.old_password, "new_password": body.new_password})(), + type( + "ChangePasswordInput", + (), + {"old_password": body.old_password, "new_password": body.new_password}, + )(), ) return Response.success(None, message="密码修改成功") @@ -220,13 +263,19 @@ def reset_password(body: ResetPasswordRequest, session: Session = Depends(get_se """邮箱+验证码重置密码(忘记密码)。""" service.reset_password( session, - ResetPasswordInput(email=body.email, code=body.code, new_password=body.new_password), + ResetPasswordInput( + email=body.email, code=body.code, new_password=body.new_password + ), ) return Response.success(None, message="密码重置成功") @router.patch("/profile", response_model=Response[UserOut]) -def update_nickname(body: UpdateNicknameRequest, request: Request, session: Session = Depends(get_session)): +def update_nickname( + body: UpdateNicknameRequest, + request: Request, + session: Session = Depends(get_session), +): """修改当前用户昵称。""" current_user = request.state.current_user user_view = service.update_nickname( @@ -237,7 +286,9 @@ def update_nickname(body: UpdateNicknameRequest, request: Request, session: Sess id=user_view.id, email=user_view.email, nickname=user_view.nickname, - email_verified_at=user_view.email_verified_at.isoformat() if user_view.email_verified_at else None, + email_verified_at=user_view.email_verified_at.isoformat() + if user_view.email_verified_at + else None, status=user_view.status, ), message="昵称修改成功", diff --git a/backend/packages/app/src/windup_app/web/api/quota.py b/backend/packages/app/src/windup_app/web/api/quota.py index 90d3f56c..85e6c257 100644 --- a/backend/packages/app/src/windup_app/web/api/quota.py +++ b/backend/packages/app/src/windup_app/web/api/quota.py @@ -4,11 +4,8 @@ -------- GET /quota/balance 查询积分余额 GET /quota/transactions 查询积分流水(分页) - -暂不实现: -POST /quota/invite/redeem 兑换邀请码 GET /quota/invite/code 获取我的邀请码 -POST /quota/invite/generate 生成新邀请码 +POST /quota/invite/generate 签发新邀请码 """ from __future__ import annotations @@ -63,15 +60,14 @@ class CreditTransactionOut(BaseModel): create_at: datetime -# -- 暂不实现 ---------------------------------------------------------------- -# -# class InviteCodeOut(BaseModel): -# """邀请码响应。""" -# ... -# -# class RedeemRequest(BaseModel): -# """兑换邀请码请求。""" -# ... +class InviteCodeOut(BaseModel): + """邀请码响应。""" + + code: str + used_count: int + expires_at: datetime + create_at: datetime + update_at: datetime # -- 端点 ---------------------------------------------------------------- @@ -102,7 +98,9 @@ def list_transactions( ) -> ListResponse[CreditTransactionOut]: """查询积分流水(分页)。""" user_id = request.state.current_user.id - txns, total = service.list_transactions(session, user_id, page=page, page_size=page_size) + txns, total = service.list_transactions( + session, user_id, page=page, page_size=page_size + ) return ListResponse.success( [CreditTransactionOut.model_validate(t) for t in txns], total=total, @@ -111,13 +109,38 @@ def list_transactions( ) -# -- 邀请码端点(暂不实现)-------------------------------------------------- -# -# @router.get("/invite/code") -# def get_invite_code(...): ... -# -# @router.post("/invite/generate") -# def generate_invite_code(...): ... -# -# @router.post("/invite/redeem") -# def redeem_invite_code(...): ... +@router.get("/invite/code", response_model=Response[InviteCodeOut]) +def get_invite_code( + request: Request, + session: Session = Depends(get_session), +) -> Response[InviteCodeOut]: + """获取当前用户未过期邀请码;没有或已过期则签发新行。""" + view = service.get_invite_code(session, request.state.current_user.id) + return Response.success( + InviteCodeOut( + code=view.code, + used_count=view.used_count, + expires_at=view.expires_at, + create_at=view.create_at, + update_at=view.update_at, + ) + ) + + +@router.post("/invite/generate", response_model=Response[InviteCodeOut]) +def generate_invite_code( + request: Request, + session: Session = Depends(get_session), +) -> Response[InviteCodeOut]: + """签发新邀请码。旧码立即过期,行保留。""" + view = service.generate_invite_code(session, request.state.current_user.id) + return Response.success( + InviteCodeOut( + code=view.code, + used_count=view.used_count, + expires_at=view.expires_at, + create_at=view.create_at, + update_at=view.update_at, + ), + message="邀请码已更新", + ) diff --git a/backend/packages/framework/src/windup_framework/config/quota.py b/backend/packages/framework/src/windup_framework/config/quota.py index 3d648feb..af1ca72c 100644 --- a/backend/packages/framework/src/windup_framework/config/quota.py +++ b/backend/packages/framework/src/windup_framework/config/quota.py @@ -9,7 +9,7 @@ class QuotaSettings(BaseSettings): """积分定价配置。 - 环境变量前缀 ``QUOTA_``,例如 ``QUOTA_REGISTER_GIFT_AMOUNT=100``。 + 环境变量前缀 ``QUOTA_``,例如 ``QUOTA_REGISTER_GIFT_AMOUNT=300``。 """ model_config = SettingsConfigDict( @@ -20,8 +20,10 @@ class QuotaSettings(BaseSettings): ) # -- 注册 / 邀请 ------------------------------------------------------- - register_gift_amount: int = 100 # 注册赠送积分 - invite_reward_amount: int = 50 # 邀请奖励(双方各得) + register_gift_amount: int = 300 # 注册赠送积分 + invite_reward_amount: int = 200 # 邀请奖励(双方各得) + invite_reward_daily_limit: int = 3 # 邀请人每日可获奖励的邀请人数(3×200=600) + invite_code_ttl_days: int = 30 # 邀请码有效期(天) # -- 生成任务 ----------------------------------------------------------- generate_image_cost: int = 10 # 生成角色参考图 diff --git a/backend/tests/conftest.py b/backend/tests/conftest.py index 10ecb720..714cbfba 100644 --- a/backend/tests/conftest.py +++ b/backend/tests/conftest.py @@ -22,7 +22,12 @@ from windup_app.bootstrap.app import create_app from windup_app.server.character.model import Character from windup_app.server.project.model import Project -from windup_app.server.quota.model import CreditAccount, CreditTransaction +from windup_app.server.quota.model import ( + CreditAccount, + CreditTransaction, + InviteCode, + InviteRecord, +) from windup_app.server.user.model import User from windup_app.server.orchestrator.model import GenerationTaskRecord from windup_app.server.workflow_run.model import WorkflowRun @@ -53,7 +58,19 @@ def _disable_generation_execution(app): app.state.run_image_task = lambda *args: None -def seed_credit_account(session, user_id: int, *, balance: int | None = None) -> CreditAccount: +def seed_invite_code(session, code: str = "AB23CD45") -> str: + """预置一个可重复使用的邀请码,供注册测试使用。""" + inviter = User(email=f"inviter-{code.lower()}@example.com", password_hash="x") + session.add(inviter) + session.flush() + session.add(InviteCode(user_id=inviter.id, code=code, used_count=0)) + seed_credit_account(session, inviter.id) + return code + + +def seed_credit_account( + session, user_id: int, *, balance: int | None = None +) -> CreditAccount: """给测试用户补一张积分账户(注册赠送口径)。""" gift = quota_settings.register_gift_amount account = CreditAccount( @@ -85,15 +102,29 @@ def _enable_sqlite_foreign_keys(dbapi_connection, _connection_record): return engine +@pytest.fixture() +def invite_code(db_session): + return seed_invite_code(db_session) + + @pytest.fixture() def engine(): """建好 ``windup_project`` 和 ``windup_user`` 表的内存 engine。""" engine = _make_engine() - Base.metadata.create_all(engine, tables=[ - Project.__table__, User.__table__, Character.__table__, WorkflowRun.__table__, - CreditAccount.__table__, CreditTransaction.__table__, - GenerationTaskRecord.__table__, - ]) + Base.metadata.create_all( + engine, + tables=[ + Project.__table__, + User.__table__, + Character.__table__, + WorkflowRun.__table__, + CreditAccount.__table__, + CreditTransaction.__table__, + InviteCode.__table__, + InviteRecord.__table__, + GenerationTaskRecord.__table__, + ], + ) yield engine engine.dispose() diff --git a/backend/tests/test_auth_api.py b/backend/tests/test_auth_api.py index 16094461..1eb8cff9 100644 --- a/backend/tests/test_auth_api.py +++ b/backend/tests/test_auth_api.py @@ -4,6 +4,8 @@ import pytest +from conftest import seed_invite_code + from windup_app.server.user.model import User from windup_app.server.user.service import _hash_password, service @@ -91,6 +93,62 @@ def test_reset_password_endpoint(auth_client, seeded_user, mock_user_redis): assert body["message"] == "密码重置成功" +def test_register_endpoint_success(client, db_session, mock_user_redis): + seed_invite_code(db_session) + db_session.commit() + mock_user_redis.get.return_value = "123456" + + resp = client.post( + "/auth/register", + json={ + "email": "invitee@example.com", + "password": "password123", + "code": "123456", + "invite_code": "AB23CD45", + "nickname": "受邀用户", + }, + ) + assert resp.status_code == 200 + body = resp.json() + assert body["code"] == 200 + assert body["message"] == "注册成功" + assert body["data"]["user"]["email"] == "invitee@example.com" + assert body["data"]["access_token"] + + +def test_register_endpoint_success_without_invite_code(client, mock_user_redis): + mock_user_redis.get.return_value = "123456" + resp = client.post( + "/auth/register", + json={ + "email": "open@example.com", + "password": "password123", + "code": "123456", + }, + ) + assert resp.status_code == 200 + body = resp.json() + assert body["code"] == 200 + assert body["data"]["user"]["email"] == "open@example.com" + assert body["data"]["access_token"] + + +def test_login_by_code_endpoint_creates_unknown_email(client, db_session, mock_user_redis): + mock_user_redis.get.return_value = "123456" + resp = client.post( + "/auth/login-by-code", + json={"email": "fresh@example.com", "code": "123456"}, + ) + assert resp.status_code == 200 + body = resp.json() + assert body["code"] == 200 + assert body["data"]["user"]["email"] == "fresh@example.com" + assert ( + db_session.query(User).filter(User.email == "fresh@example.com").one_or_none() + is not None + ) + + def test_update_nickname_endpoint(auth_client, seeded_user, mock_user_redis): resp = auth_client.patch("/auth/profile", json={"nickname": "新昵称"}) assert resp.status_code == 200 diff --git a/backend/tests/test_auth_registration_closed.py b/backend/tests/test_auth_registration_closed.py index aa77c335..28fa99df 100644 --- a/backend/tests/test_auth_registration_closed.py +++ b/backend/tests/test_auth_registration_closed.py @@ -1,30 +1,25 @@ -"""内测关闭公开注册:公开建号路径必须被拒绝。""" +"""无效邀请码不得建号。""" from windup_common.enums.biz_code import BizCode +from windup_app.server.user.model import User -def test_register_endpoint_rejects_public_signup(client): + +def test_register_endpoint_rejects_invalid_invite_code(client, db_session): resp = client.post( "/auth/register", json={ "email": "new@example.com", "password": "password123", "code": "123456", + "invite_code": "NOPE1234", }, ) assert resp.status_code == 200 body = resp.json() assert body["code"] == BizCode.BAD_REQUEST - assert body["message"] == "内测期间暂不开放注册" - assert body["data"] is None - - -def test_send_code_rejects_register_purpose(client): - resp = client.post( - "/auth/send-code", - json={"email": "new@example.com", "purpose": "register"}, + assert body["message"] == "邀请码无效" + assert ( + db_session.query(User).filter(User.email == "new@example.com").one_or_none() + is None ) - assert resp.status_code == 200 - body = resp.json() - assert body["code"] == BizCode.BAD_REQUEST - assert body["message"] == "内测期间暂不开放注册" diff --git a/backend/tests/test_quota.py b/backend/tests/test_quota.py index 64364464..775c8176 100644 --- a/backend/tests/test_quota.py +++ b/backend/tests/test_quota.py @@ -106,18 +106,25 @@ def test_get_nonexistent_account(self, db_session, quota_service): class TestReserveCredit: def test_reserve_success(self, db_session, quota_service, user_with_account): uid = user_with_account.id - quota_service.reserve_credit(db_session, uid, quota_settings.generate_image_cost, "task:1") + quota_service.reserve_credit( + db_session, uid, quota_settings.generate_image_cost, "task:1" + ) account = db_session.scalar( select(CreditAccount).where(CreditAccount.user_id == uid) ) - assert account.balance == quota_settings.register_gift_amount - quota_settings.generate_image_cost + assert ( + account.balance + == quota_settings.register_gift_amount - quota_settings.generate_image_cost + ) assert account.frozen == quota_settings.generate_image_cost def test_reserve_insufficient(self, db_session, quota_service, user_with_account): uid = user_with_account.id with pytest.raises(BizException, match="积分不足"): - quota_service.reserve_credit(db_session, uid, quota_settings.register_gift_amount + 1, "task:2") + quota_service.reserve_credit( + db_session, uid, quota_settings.register_gift_amount + 1, "task:2" + ) def test_reserve_nonexistent_account(self, db_session, quota_service): with pytest.raises(BizException, match="积分账户不存在"): @@ -125,11 +132,14 @@ def test_reserve_nonexistent_account(self, db_session, quota_service): def test_reserve_writes_txn(self, db_session, quota_service, user_with_account): uid = user_with_account.id - quota_service.reserve_credit(db_session, uid, quota_settings.generate_image_cost, "task:3") + quota_service.reserve_credit( + db_session, uid, quota_settings.generate_image_cost, "task:3" + ) txn = db_session.scalar( - select(CreditTransaction) - .where(CreditTransaction.user_id == uid, CreditTransaction.ref_id == "task:3") + select(CreditTransaction).where( + CreditTransaction.user_id == uid, CreditTransaction.ref_id == "task:3" + ) ) assert txn is not None assert txn.delta == -quota_settings.generate_image_cost @@ -157,8 +167,12 @@ def test_capture_full(self, db_session, quota_service, user_with_account): def test_capture_partial_refund(self, db_session, quota_service, user_with_account): """冻结 50,实际扣 30,差额 20 退回。""" uid = user_with_account.id - quota_service.reserve_credit(db_session, uid, quota_settings.generate_action_cost, "task:4") - quota_service.capture_credit(db_session, uid, 30, "task:4", quota_settings.generate_action_cost) + quota_service.reserve_credit( + db_session, uid, quota_settings.generate_action_cost, "task:4" + ) + quota_service.capture_credit( + db_session, uid, 30, "task:4", quota_settings.generate_action_cost + ) account = db_session.scalar( select(CreditAccount).where(CreditAccount.user_id == uid) @@ -167,11 +181,17 @@ def test_capture_partial_refund(self, db_session, quota_service, user_with_accou assert account.frozen == 0 assert account.total_spent == 30 - def test_capture_writes_txn_and_refund(self, db_session, quota_service, user_with_account): + def test_capture_writes_txn_and_refund( + self, db_session, quota_service, user_with_account + ): """有差额退回时应写两条流水:扣减 + 退款。""" uid = user_with_account.id - quota_service.reserve_credit(db_session, uid, quota_settings.generate_action_cost, "task:5") - quota_service.capture_credit(db_session, uid, 30, "task:5", quota_settings.generate_action_cost) + quota_service.reserve_credit( + db_session, uid, quota_settings.generate_action_cost, "task:5" + ) + quota_service.capture_credit( + db_session, uid, 30, "task:5", quota_settings.generate_action_cost + ) txns = db_session.scalars( select(CreditTransaction).where(CreditTransaction.user_id == uid) @@ -180,10 +200,14 @@ def test_capture_writes_txn_and_refund(self, db_session, quota_service, user_wit assert CreditReason.CAPTURED in reasons assert CreditReason.REFUND in reasons - def test_capture_insufficient_frozen(self, db_session, quota_service, user_with_account): + def test_capture_insufficient_frozen( + self, db_session, quota_service, user_with_account + ): """冻结额度不足时应抛异常。""" uid = user_with_account.id - quota_service.reserve_credit(db_session, uid, quota_settings.generate_image_cost, "task:6") + quota_service.reserve_credit( + db_session, uid, quota_settings.generate_image_cost, "task:6" + ) with pytest.raises(BizException, match="冻结额度不足"): quota_service.capture_credit(db_session, uid, 100, "task:6", 100) @@ -198,8 +222,12 @@ def test_capture_nonexistent_account(self, db_session, quota_service): class TestReleaseCredit: def test_release_success(self, db_session, quota_service, user_with_account): uid = user_with_account.id - quota_service.reserve_credit(db_session, uid, quota_settings.generate_image_cost, "task:7") - quota_service.release_credit(db_session, uid, quota_settings.generate_image_cost, "task:7") + quota_service.reserve_credit( + db_session, uid, quota_settings.generate_image_cost, "task:7" + ) + quota_service.release_credit( + db_session, uid, quota_settings.generate_image_cost, "task:7" + ) account = db_session.scalar( select(CreditAccount).where(CreditAccount.user_id == uid) @@ -209,17 +237,25 @@ def test_release_success(self, db_session, quota_service, user_with_account): def test_release_writes_txn(self, db_session, quota_service, user_with_account): uid = user_with_account.id - quota_service.reserve_credit(db_session, uid, quota_settings.generate_image_cost, "task:8") - quota_service.release_credit(db_session, uid, quota_settings.generate_image_cost, "task:8") + quota_service.reserve_credit( + db_session, uid, quota_settings.generate_image_cost, "task:8" + ) + quota_service.release_credit( + db_session, uid, quota_settings.generate_image_cost, "task:8" + ) txn = db_session.scalar( - select(CreditTransaction) - .where(CreditTransaction.user_id == uid, CreditTransaction.reason == CreditReason.REFUND) + select(CreditTransaction).where( + CreditTransaction.user_id == uid, + CreditTransaction.reason == CreditReason.REFUND, + ) ) assert txn is not None assert txn.delta == quota_settings.generate_image_cost - def test_release_insufficient_frozen(self, db_session, quota_service, user_with_account): + def test_release_insufficient_frozen( + self, db_session, quota_service, user_with_account + ): uid = user_with_account.id with pytest.raises(BizException, match="冻结额度不足"): quota_service.release_credit(db_session, uid, 100, "task:9") @@ -248,8 +284,9 @@ def test_credit_writes_txn(self, db_session, quota_service, user_with_account): quota_service.credit(db_session, uid, 50, CreditReason.ADMIN_ADJUST, "admin:2") txn = db_session.scalar( - select(CreditTransaction) - .where(CreditTransaction.user_id == uid, CreditTransaction.ref_id == "admin:2") + select(CreditTransaction).where( + CreditTransaction.user_id == uid, CreditTransaction.ref_id == "admin:2" + ) ) assert txn is not None assert txn.delta == 50 @@ -290,8 +327,16 @@ def test_list_empty(self, db_session, quota_service, user_with_account): def test_list_after_operations(self, db_session, quota_service, user_with_account): uid = user_with_account.id - quota_service.reserve_credit(db_session, uid, quota_settings.generate_image_cost, "task:10") - quota_service.capture_credit(db_session, uid, quota_settings.generate_image_cost, "task:10", quota_settings.generate_image_cost) + quota_service.reserve_credit( + db_session, uid, quota_settings.generate_image_cost, "task:10" + ) + quota_service.capture_credit( + db_session, + uid, + quota_settings.generate_image_cost, + "task:10", + quota_settings.generate_image_cost, + ) txns, total = quota_service.list_transactions(db_session, uid) assert total >= 2 @@ -300,13 +345,19 @@ def test_list_after_operations(self, db_session, quota_service, user_with_accoun def test_list_pagination(self, db_session, quota_service, user_with_account): uid = user_with_account.id for i in range(5): - quota_service.credit(db_session, uid, 10, CreditReason.ADMIN_ADJUST, f"page:{i}") + quota_service.credit( + db_session, uid, 10, CreditReason.ADMIN_ADJUST, f"page:{i}" + ) - txns_p1, total = quota_service.list_transactions(db_session, uid, page=1, page_size=2) + txns_p1, total = quota_service.list_transactions( + db_session, uid, page=1, page_size=2 + ) assert total == 5 assert len(txns_p1) == 2 - txns_p3, _ = quota_service.list_transactions(db_session, uid, page=3, page_size=2) + txns_p3, _ = quota_service.list_transactions( + db_session, uid, page=3, page_size=2 + ) assert len(txns_p3) == 1 # 最后一页只有 1 条 def test_list_other_user_empty(self, db_session, quota_service, user_with_account): @@ -409,7 +460,9 @@ def test_list_transactions_empty(self, auth_quota_client): assert data["data"] == [] assert data["total"] == 0 - def test_list_transactions_pagination(self, auth_quota_client, db_session, user_with_account): + def test_list_transactions_pagination( + self, auth_quota_client, db_session, user_with_account + ): """先写入几条流水,再通过 API 分页查询。""" uid = user_with_account.id service = SqlAlchemyQuotaService() @@ -423,7 +476,9 @@ def test_list_transactions_pagination(self, auth_quota_client, db_session, user_ assert data["total"] == 5 assert len(data["data"]) == 2 - def test_list_transactions_default_pagination(self, auth_quota_client, db_session, user_with_account): + def test_list_transactions_default_pagination( + self, auth_quota_client, db_session, user_with_account + ): """默认分页参数。""" uid = user_with_account.id service = SqlAlchemyQuotaService() @@ -441,3 +496,346 @@ def test_unauthenticated_access(self, client): assert resp.status_code == 200 data = resp.json() assert data["code"] == 401 + + +def _gift_account(session: Session, user_id: int) -> None: + session.add( + CreditAccount( + user_id=user_id, + balance=quota_settings.register_gift_amount, + frozen=0, + total_earned=quota_settings.register_gift_amount, + total_spent=0, + ) + ) + session.flush() + + +class TestInviteCode: + """邀请码生成、查询与兑换。""" + + def test_get_invite_code_creates_when_missing(self, auth_quota_client): + resp = auth_quota_client.get("/quota/invite/code") + assert resp.status_code == 200 + data = resp.json() + assert data["code"] == 200 + assert len(data["data"]["code"]) == 8 + assert data["data"]["used_count"] == 0 + + again = auth_quota_client.get("/quota/invite/code") + assert again.json()["data"]["code"] == data["data"]["code"] + assert again.json()["data"]["expires_at"] + + def test_generate_invite_code_rotates(self, auth_quota_client): + first = auth_quota_client.get("/quota/invite/code").json()["data"]["code"] + second = auth_quota_client.post("/quota/invite/generate").json()["data"]["code"] + assert second != first + assert len(second) == 8 + assert auth_quota_client.get("/quota/invite/code").json()["data"]["code"] == second + + def test_generate_invite_code_locks_existing_row( + self, db_session, quota_service, monkeypatch + ): + from sqlalchemy.sql.selectable import Select + from windup_app.server.user.model import User + + host = User(email="lock-host@example.com", password_hash="x") + db_session.add(host) + db_session.flush() + quota_service.generate_invite_code(db_session, host.id) + + locked = [] + original = Select.with_for_update + + def tracking(self, *args, **kwargs): + locked.append(True) + return original(self, *args, **kwargs) + + monkeypatch.setattr(Select, "with_for_update", tracking) + quota_service.generate_invite_code(db_session, host.id) + assert locked, "轮换已有邀请码时应对该行加 FOR UPDATE" + + def test_generate_invite_code_keeps_old_row(self, db_session, quota_service): + from datetime import datetime, timezone + from sqlalchemy import select + from windup_app.server.quota.model import InviteCode + from windup_app.server.user.model import User + + host = User(email="append-host@example.com", password_hash="x") + db_session.add(host) + db_session.flush() + first = quota_service.generate_invite_code(db_session, host.id) + second = quota_service.generate_invite_code(db_session, host.id) + rows = db_session.scalars( + select(InviteCode).where(InviteCode.user_id == host.id) + ).all() + assert {row.code for row in rows} == {first.code, second.code} + old = next(row for row in rows if row.code == first.code) + now = datetime.now(timezone.utc) + exp = old.expires_at if old.expires_at.tzinfo else old.expires_at.replace( + tzinfo=timezone.utc + ) + assert exp <= now + + def test_get_invite_code_issues_new_row_after_expiry( + self, db_session, quota_service + ): + from datetime import datetime, timedelta, timezone + from sqlalchemy import select + from windup_app.server.quota.model import InviteCode + from windup_app.server.user.model import User + + host = User(email="expire-host@example.com", password_hash="x") + db_session.add(host) + db_session.flush() + first = quota_service.generate_invite_code(db_session, host.id) + row = db_session.scalar( + select(InviteCode).where(InviteCode.code == first.code) + ) + row.expires_at = datetime.now(timezone.utc) - timedelta(seconds=1) + db_session.flush() + + second = quota_service.get_invite_code(db_session, host.id) + assert second.code != first.code + assert ( + db_session.scalar( + select(InviteCode).where(InviteCode.code == first.code) + ) + is not None + ) + + def test_redeem_unique_violation_is_already_redeemed( + self, db_session, quota_service, monkeypatch + ): + """并发双兑时 unique(invitee_id) 应收敛为「已填写过邀请码」,而不是 500。""" + from sqlalchemy.exc import IntegrityError + from windup_app.server.user.model import User + from windup_common.exceptions import BizException + + host = User(email="race-host@example.com", password_hash="x") + guest = User(email="race-guest@example.com", password_hash="x") + db_session.add_all([host, guest]) + db_session.flush() + _gift_account(db_session, host.id) + _gift_account(db_session, guest.id) + view = quota_service.generate_invite_code(db_session, host.id) + + from windup_app.server.quota.model import InviteRecord + + orig_flush = db_session.flush + + def boom(*_args, **_kwargs): + if any(isinstance(obj, InviteRecord) for obj in db_session.new): + raise IntegrityError( + "INSERT", + {}, + Exception( + "UNIQUE constraint failed: windup_invite_record.invitee_id" + ), + ) + return orig_flush(*_args, **_kwargs) + + monkeypatch.setattr(db_session, "flush", boom) + + with pytest.raises(BizException, match="已填写过邀请码"): + quota_service.redeem_invite_code(db_session, guest.id, view.code) + + def test_redeem_invite_code_rewards_both_users(self, db_session, quota_service): + from windup_app.server.user.model import User + + inviter = User(email="host@example.com", password_hash="x") + invitee = User(email="guest@example.com", password_hash="x") + db_session.add_all([inviter, invitee]) + db_session.flush() + _gift_account(db_session, inviter.id) + _gift_account(db_session, invitee.id) + view = quota_service.generate_invite_code(db_session, inviter.id) + + quota_service.redeem_invite_code(db_session, invitee.id, view.code.lower()) + + host = quota_service.get_account(db_session, inviter.id) + guest = quota_service.get_account(db_session, invitee.id) + assert ( + host.balance + == quota_settings.register_gift_amount + quota_settings.invite_reward_amount + ) + assert ( + guest.balance + == quota_settings.register_gift_amount + quota_settings.invite_reward_amount + ) + + def test_inviter_daily_reward_stops_after_three_invites( + self, db_session, quota_service + ): + from windup_app.server.quota.model import InviteRecord + from windup_app.server.user.model import User + + inviter = User(email="cap-host@example.com", password_hash="x") + db_session.add(inviter) + db_session.flush() + _gift_account(db_session, inviter.id) + view = quota_service.generate_invite_code(db_session, inviter.id) + + guests = [] + for i in range(4): + guest = User(email=f"cap-guest-{i}@example.com", password_hash="x") + db_session.add(guest) + db_session.flush() + _gift_account(db_session, guest.id) + quota_service.redeem_invite_code(db_session, guest.id, view.code) + guests.append(guest) + + host = quota_service.get_account(db_session, inviter.id) + assert host.balance == quota_settings.register_gift_amount + ( + quota_settings.invite_reward_amount * 3 + ) + assert ( + db_session.scalar( + select(InviteRecord.id).where( + InviteRecord.invitee_id == guests[3].id + ) + ) + is not None + ) + fourth = quota_service.get_account(db_session, guests[3].id) + assert ( + fourth.balance + == quota_settings.register_gift_amount + quota_settings.invite_reward_amount + ) + + def test_inviter_daily_reward_resets_next_utc_day( + self, db_session, quota_service + ): + from datetime import timedelta + from windup_app.server.quota.model import InviteRecord + from windup_app.server.quota.service import _now + from windup_app.server.user.model import User + + inviter = User(email="nextday-host@example.com", password_hash="x") + db_session.add(inviter) + db_session.flush() + _gift_account(db_session, inviter.id) + view = quota_service.generate_invite_code(db_session, inviter.id) + + for i in range(3): + guest = User(email=f"old-guest-{i}@example.com", password_hash="x") + db_session.add(guest) + db_session.flush() + _gift_account(db_session, guest.id) + quota_service.redeem_invite_code(db_session, guest.id, view.code) + + yesterday = _now() - timedelta(days=1) + for row in db_session.scalars( + select(InviteRecord).where(InviteRecord.inviter_id == inviter.id) + ).all(): + row.create_at = yesterday + db_session.flush() + + today_guest = User(email="today-guest@example.com", password_hash="x") + db_session.add(today_guest) + db_session.flush() + _gift_account(db_session, today_guest.id) + quota_service.redeem_invite_code(db_session, today_guest.id, view.code) + + host = quota_service.get_account(db_session, inviter.id) + assert host.balance == quota_settings.register_gift_amount + ( + quota_settings.invite_reward_amount * 4 + ) + + def test_redeem_rejects_own_code_and_repeat(self, db_session, quota_service): + from windup_app.server.user.model import User + from windup_common.exceptions import BizException + + host = User(email="self@example.com", password_hash="x") + guest = User(email="once@example.com", password_hash="x") + db_session.add_all([host, guest]) + db_session.flush() + _gift_account(db_session, host.id) + _gift_account(db_session, guest.id) + view = quota_service.generate_invite_code(db_session, host.id) + + with pytest.raises(BizException, match="不能填写自己的邀请码"): + quota_service.redeem_invite_code(db_session, host.id, view.code) + + quota_service.redeem_invite_code(db_session, guest.id, view.code) + with pytest.raises(BizException, match="已填写过邀请码"): + quota_service.redeem_invite_code(db_session, guest.id, view.code) + + def test_generate_invite_code_rejects_missing_user(self, db_session, quota_service): + from windup_common.exceptions import BizException + + with pytest.raises(BizException, match="用户不存在"): + quota_service.generate_invite_code(db_session, 999999) + + def test_allocate_invite_code_gives_up_on_collision( + self, db_session, quota_service, monkeypatch + ): + from windup_app.server.quota import service as quota_mod + from windup_app.server.user.model import User + from windup_common.exceptions import BizException + + taken = User(email="taken@example.com", password_hash="x") + host = User(email="alloc@example.com", password_hash="x") + db_session.add_all([taken, host]) + db_session.flush() + occupied = quota_service.generate_invite_code(db_session, taken.id) + monkeypatch.setattr(quota_mod, "_new_invite_code", lambda: occupied.code) + + with pytest.raises(BizException, match="邀请码生成失败"): + quota_service.generate_invite_code(db_session, host.id) + + def test_redeem_rejects_blank_or_unknown_code(self, db_session, quota_service): + from windup_app.server.user.model import User + from windup_common.exceptions import BizException + + guest = User(email="blank@example.com", password_hash="x") + db_session.add(guest) + db_session.flush() + _gift_account(db_session, guest.id) + + with pytest.raises(BizException, match="邀请码无效"): + quota_service.redeem_invite_code(db_session, guest.id, " ") + with pytest.raises(BizException, match="邀请码无效"): + quota_service.redeem_invite_code(db_session, guest.id, "IO01") + with pytest.raises(BizException, match="邀请码无效"): + quota_service.redeem_invite_code(db_session, guest.id, "NOPE1234") + + def test_redeem_rejects_expired_code(self, db_session, quota_service): + from datetime import datetime, timedelta, timezone + from sqlalchemy import select + from windup_app.server.quota.model import InviteCode + from windup_app.server.user.model import User + from windup_common.enums.biz_code import BizCode + from windup_common.exceptions import BizException + + host = User(email="stale-host@example.com", password_hash="x") + guest = User(email="stale-guest@example.com", password_hash="x") + db_session.add_all([host, guest]) + db_session.flush() + _gift_account(db_session, host.id) + _gift_account(db_session, guest.id) + view = quota_service.generate_invite_code(db_session, host.id) + row = db_session.scalar(select(InviteCode).where(InviteCode.code == view.code)) + row.expires_at = datetime.now(timezone.utc) - timedelta(seconds=1) + db_session.flush() + + with pytest.raises(BizException, match="邀请码已过期") as exc: + quota_service.redeem_invite_code(db_session, guest.id, view.code) + assert exc.value.code == BizCode.NOT_FOUND + + def test_redeem_rejects_missing_invitee(self, db_session, quota_service): + from windup_app.server.user.model import User + from windup_common.exceptions import BizException + + host = User(email="orphan-host@example.com", password_hash="x") + db_session.add(host) + db_session.flush() + _gift_account(db_session, host.id) + view = quota_service.generate_invite_code(db_session, host.id) + + with pytest.raises(BizException, match="用户不存在"): + quota_service.redeem_invite_code(db_session, 999999, view.code) + + def test_invite_redeem_endpoint_removed(self, auth_quota_client): + resp = auth_quota_client.post("/quota/invite/redeem", json={"code": "AB23CD45"}) + assert resp.status_code == 404 diff --git a/backend/tests/test_user_service.py b/backend/tests/test_user_service.py index 671a677a..31cc1761 100644 --- a/backend/tests/test_user_service.py +++ b/backend/tests/test_user_service.py @@ -31,6 +31,12 @@ # -- Fixtures ------------------------------------------------------------ +@pytest.fixture(autouse=True) +def _seed_invite(request): + if "db_session" in request.fixturenames: + request.getfixturevalue("invite_code") + + @pytest.fixture() def mock_redis(): """Mock Redis 客户端。""" @@ -129,6 +135,7 @@ def test_register_success(db_session, service, mock_email): email="new@example.com", password="password123", code="123456", + invite_code="AB23CD45", ) result = service.register_by_email(db_session, input_data) @@ -148,6 +155,7 @@ def test_public_methods_accept_session(db_session, service, mock_email): email="public@example.com", password="password123", code="123456", + invite_code="AB23CD45", ), ) @@ -191,19 +199,21 @@ def test_register_creates_credit_account(db_session, service, mock_email): email="credit@example.com", password="password123", code="123456", + invite_code="AB23CD45", ) result = service.register_by_email(db_session, input_data) user_id = result.user.id + expected = quota_settings.register_gift_amount + quota_settings.invite_reward_amount # 验证积分账户已创建 account = db_session.scalar( select(CreditAccount).where(CreditAccount.user_id == user_id) ) assert account is not None - assert account.balance == quota_settings.register_gift_amount + assert account.balance == expected assert account.frozen == 0 - assert account.total_earned == quota_settings.register_gift_amount + assert account.total_earned == expected assert account.total_spent == 0 # 验证赠送流水已记录 @@ -222,7 +232,12 @@ def test_register_creates_credit_account(db_session, service, mock_email): def test_register_duplicate_email(db_session, service): # 先注册一个用户 service._redis.get.return_value = "123456" - input_data = RegisterInput(email="dup@example.com", password="pass123", code="123456") + input_data = RegisterInput( + email="dup@example.com", + password="pass123", + code="123456", + invite_code="AB23CD45", + ) service.register_by_email(db_session, input_data) # 尝试重复注册 @@ -237,6 +252,7 @@ def test_register_wrong_code(db_session, service): email="new@example.com", password="password123", code="999999", # 错误验证码 + invite_code="AB23CD45", ) with pytest.raises(BizException, match="验证码错误"): @@ -250,19 +266,84 @@ def test_register_expired_code(db_session, service): email="new@example.com", password="password123", code="123456", + invite_code="AB23CD45", ) with pytest.raises(BizException, match="验证码已过期"): service.register_by_email(db_session, input_data) +def test_register_blank_invite_code_only_gives_register_gift(db_session, service): + """未带邀请码时只发注册赠送,不挡注册。""" + from sqlalchemy import select + from windup_app.server.quota.model import CreditAccount + from windup_framework.config.quota import settings as quota_settings + + service._redis.get.return_value = "123456" + input_data = RegisterInput( + email="blank-invite@example.com", + password="password123", + code="123456", + invite_code=" ", + ) + + result = service.register_by_email(db_session, input_data) + account = db_session.scalar( + select(CreditAccount).where(CreditAccount.user_id == result.user.id) + ) + assert account is not None + assert account.balance == quota_settings.register_gift_amount + + +def test_register_rejects_invite_code_outside_link_charset(db_session, service): + """前端邀请链接用 A-H/J-N/P-Z/2-9,含 I/O/0/1 的码不会进注册请求。""" + service._redis.get.return_value = "123456" + input_data = RegisterInput( + email="bad-charset@example.com", + password="password123", + code="123456", + invite_code="IIII", + ) + + with pytest.raises(BizException, match="邀请码无效"): + service.register_by_email(db_session, input_data) + + +def test_register_expired_invite_code(db_session, service): + from datetime import datetime, timedelta, timezone + from sqlalchemy import select + from windup_app.server.quota.model import InviteCode + from windup_common.enums.biz_code import BizCode + + row = db_session.scalar(select(InviteCode).where(InviteCode.code == "AB23CD45")) + row.expires_at = datetime.now(timezone.utc) - timedelta(days=1) + db_session.flush() + + service._redis.get.return_value = "123456" + input_data = RegisterInput( + email="late@example.com", + password="password123", + code="123456", + invite_code="AB23CD45", + ) + with pytest.raises(BizException, match="邀请码已过期") as exc: + service.register_by_email(db_session, input_data) + assert exc.value.code == BizCode.NOT_FOUND + assert db_session.scalar(select(User).where(User.email == "late@example.com")) is None + + # -- 登录测试 ------------------------------------------------------------ def test_login_success(db_session, service, mock_email): # 先注册 service._redis.get.return_value = "123456" - register_input = RegisterInput(email="login@example.com", password="pass123", code="123456") + register_input = RegisterInput( + email="login@example.com", + password="pass123", + code="123456", + invite_code="AB23CD45", + ) service.register_by_email(db_session, register_input) # 登录(不需要验证码) @@ -277,7 +358,12 @@ def test_login_success(db_session, service, mock_email): def test_login_wrong_password(db_session, service, mock_email): # 先注册 service._redis.get.return_value = "123456" - register_input = RegisterInput(email="login@example.com", password="pass123", code="123456") + register_input = RegisterInput( + email="login@example.com", + password="pass123", + code="123456", + invite_code="AB23CD45", + ) service.register_by_email(db_session, register_input) # 密码错误 @@ -301,11 +387,17 @@ def test_login_nonexistent_user(db_session, service): def test_login_banned_user(db_session, service, mock_email): # 先注册 service._redis.get.return_value = "123456" - register_input = RegisterInput(email="banned@example.com", password="pass123", code="123456") + register_input = RegisterInput( + email="banned@example.com", + password="pass123", + code="123456", + invite_code="AB23CD45", + ) service.register_by_email(db_session, register_input) # 封禁用户 from sqlalchemy import select + user = db_session.scalar(select(User).where(User.email == "banned@example.com")) user.status = UserStatus.BANNED db_session.flush() @@ -321,29 +413,36 @@ def test_login_banned_user(db_session, service, mock_email): # -- 验证码登录测试 ------------------------------------------------------ -def test_login_by_code_unknown_email_does_not_create_user(db_session, service, mock_email): - """内测关闭公开注册后,验证码登录不得自动建号。""" +def test_login_by_code_unknown_email_creates_user_and_gifts( + db_session, service, mock_email +): + """未知邮箱验证码登录自动建号,并只发注册赠送。""" from sqlalchemy import select + from windup_app.server.quota.model import CreditAccount + from windup_framework.config.quota import settings as quota_settings service._redis.get.return_value = "123456" input_data = LoginByCodeInput(email="code@example.com", code="123456") - with pytest.raises(BizException, match="账号不存在") as exc: - service.login_by_code(db_session, input_data) - - from windup_common.enums.biz_code import BizCode + result = service.login_by_code(db_session, input_data) - assert exc.value.code == BizCode.NOT_FOUND - assert db_session.scalar(select(User).where(User.email == "code@example.com")) is None + user = db_session.scalar(select(User).where(User.email == "code@example.com")) + assert user is not None + assert result.user.id == user.id + assert result.user.email_verified_at is not None + account = db_session.scalar( + select(CreditAccount).where(CreditAccount.user_id == user.id) + ) + assert account is not None + assert account.balance == quota_settings.register_gift_amount -def test_send_verification_code_rejects_register_purpose(service, mock_email): +def test_send_verification_code_allows_register_purpose(service, mock_email): service._redis.get.return_value = None - with pytest.raises(BizException, match="内测期间暂不开放注册"): - service.send_verification_code("new@example.com", "register") + service.send_verification_code("new@example.com", "register") - mock_email.send_verification_code.assert_not_called() + mock_email.send_verification_code.assert_called_once() def test_login_by_code_banned_user(db_session, service, mock_email): @@ -352,9 +451,16 @@ def test_login_by_code_banned_user(db_session, service, mock_email): service._redis.get.return_value = "123456" service.register_by_email( db_session, - RegisterInput(email="banned-code@example.com", password="pass123", code="123456"), + RegisterInput( + email="banned-code@example.com", + password="pass123", + code="123456", + invite_code="AB23CD45", + ), + ) + user = db_session.scalar( + select(User).where(User.email == "banned-code@example.com") ) - user = db_session.scalar(select(User).where(User.email == "banned-code@example.com")) user.status = UserStatus.BANNED db_session.flush() @@ -384,7 +490,12 @@ def test_login_by_code_marks_unverified_email(db_session, service): def test_login_by_code_existing_user(db_session, service, mock_email): # 先注册 service._redis.get.return_value = "123456" - register_input = RegisterInput(email="exist@example.com", password="pass123", code="123456") + register_input = RegisterInput( + email="exist@example.com", + password="pass123", + code="123456", + invite_code="AB23CD45", + ) service.register_by_email(db_session, register_input) # 验证码登录 @@ -490,16 +601,25 @@ def test_refresh_tokens_concurrent_reuse(service, mock_redis): def test_change_password(db_session, service, mock_email): # 先注册 service._redis.get.return_value = "123456" - register_input = RegisterInput(email="change@example.com", password="oldpass123", code="123456") + register_input = RegisterInput( + email="change@example.com", + password="oldpass123", + code="123456", + invite_code="AB23CD45", + ) result = service.register_by_email(db_session, register_input) # 修改密码 - change_input = ChangePasswordInput(old_password="oldpass123", new_password="newpass123") + change_input = ChangePasswordInput( + old_password="oldpass123", new_password="newpass123" + ) service.change_password(db_session, result.user.id, change_input) # 用新密码登录 service._redis.get.return_value = None - login_input = LoginByPasswordInput(email="change@example.com", password="newpass123") + login_input = LoginByPasswordInput( + email="change@example.com", password="newpass123" + ) login_result = service.login_by_password(db_session, login_input) assert login_result.user.email == "change@example.com" @@ -508,7 +628,12 @@ def test_change_password(db_session, service, mock_email): def test_change_password_wrong_old(db_session, service, mock_email): # 先注册 service._redis.get.return_value = "123456" - register_input = RegisterInput(email="change@example.com", password="oldpass123", code="123456") + register_input = RegisterInput( + email="change@example.com", + password="oldpass123", + code="123456", + invite_code="AB23CD45", + ) result = service.register_by_email(db_session, register_input) # 旧密码错误 @@ -524,7 +649,13 @@ def test_change_password_wrong_old(db_session, service, mock_email): def test_update_nickname(db_session, service, mock_email): """修改昵称后立即生效。""" service._redis.get.return_value = "123456" - register_input = RegisterInput(email="nick@example.com", password="pass1234", code="123456", nickname="旧昵称") + register_input = RegisterInput( + email="nick@example.com", + password="pass1234", + code="123456", + nickname="旧昵称", + invite_code="AB23CD45", + ) result = service.register_by_email(db_session, register_input) update_input = UpdateNicknameInput(nickname="新昵称") @@ -537,7 +668,12 @@ def test_update_nickname(db_session, service, mock_email): def test_update_nickname_max_length(db_session, service, mock_email): """昵称长度上限 50。""" service._redis.get.return_value = "123456" - register_input = RegisterInput(email="nick2@example.com", password="pass1234", code="123456") + register_input = RegisterInput( + email="nick2@example.com", + password="pass1234", + code="123456", + invite_code="AB23CD45", + ) result = service.register_by_email(db_session, register_input) long_nickname = "a" * 50 @@ -562,12 +698,19 @@ def test_reset_password(db_session, service, mock_email): """邮箱+验证码重置密码后,新密码可登录。""" # 先注册 service._redis.get.return_value = "123456" - register_input = RegisterInput(email="reset@example.com", password="oldpass123", code="123456") + register_input = RegisterInput( + email="reset@example.com", + password="oldpass123", + code="123456", + invite_code="AB23CD45", + ) service.register_by_email(db_session, register_input) # 重置密码(验证码 purpose 为 reset_password) service._redis.get.return_value = "654321" - reset_input = ResetPasswordInput(email="reset@example.com", code="654321", new_password="newpass123") + reset_input = ResetPasswordInput( + email="reset@example.com", code="654321", new_password="newpass123" + ) service.reset_password(db_session, reset_input) # 用新密码登录 @@ -581,12 +724,19 @@ def test_reset_password(db_session, service, mock_email): def test_reset_password_wrong_code(db_session, service, mock_email): """验证码错误时拒绝重置。""" service._redis.get.return_value = "123456" - register_input = RegisterInput(email="reset2@example.com", password="oldpass123", code="123456") + register_input = RegisterInput( + email="reset2@example.com", + password="oldpass123", + code="123456", + invite_code="AB23CD45", + ) service.register_by_email(db_session, register_input) # 验证码错误 service._redis.get.return_value = None # 验证码过期 - reset_input = ResetPasswordInput(email="reset2@example.com", code="000000", new_password="newpass123") + reset_input = ResetPasswordInput( + email="reset2@example.com", code="000000", new_password="newpass123" + ) with pytest.raises(BizException, match="验证码已过期"): service.reset_password(db_session, reset_input) @@ -595,7 +745,9 @@ def test_reset_password_wrong_code(db_session, service, mock_email): def test_reset_password_user_not_found(db_session, service): """用户不存在时拒绝重置。""" service._redis.get.return_value = "654321" - reset_input = ResetPasswordInput(email="noexist@example.com", code="654321", new_password="newpass123") + reset_input = ResetPasswordInput( + email="noexist@example.com", code="654321", new_password="newpass123" + ) with pytest.raises(BizException, match="用户不存在"): service.reset_password(db_session, reset_input) @@ -608,7 +760,12 @@ def test_login_account_locked(db_session, service, mock_email): """账号被锁定后拒绝登录(即使密码正确)。""" # 先注册 service._redis.get.return_value = "123456" - register_input = RegisterInput(email="lock@example.com", password="pass123", code="123456") + register_input = RegisterInput( + email="lock@example.com", + password="pass123", + code="123456", + invite_code="AB23CD45", + ) service.register_by_email(db_session, register_input) # 模拟账号锁定 diff --git a/openapi.json b/openapi.json index 0e5d0242..33cc2cde 100644 --- a/openapi.json +++ b/openapi.json @@ -720,6 +720,43 @@ "title": "HTTPValidationError", "type": "object" }, + "InviteCodeOut": { + "description": "邀请码响应。", + "properties": { + "code": { + "title": "Code", + "type": "string" + }, + "create_at": { + "format": "date-time", + "title": "Create At", + "type": "string" + }, + "expires_at": { + "format": "date-time", + "title": "Expires At", + "type": "string" + }, + "update_at": { + "format": "date-time", + "title": "Update At", + "type": "string" + }, + "used_count": { + "title": "Used Count", + "type": "integer" + } + }, + "required": [ + "code", + "used_count", + "expires_at", + "create_at", + "update_at" + ], + "title": "InviteCodeOut", + "type": "object" + }, "ListResponse_CharacterOut_": { "properties": { "code": { @@ -1242,6 +1279,19 @@ "title": "Email", "type": "string" }, + "invite_code": { + "anyOf": [ + { + "maxLength": 16, + "type": "string" + }, + { + "type": "null" + } + ], + "description": "邀请链接中的邀请码,选填;有则发双方邀请奖励", + "title": "Invite Code" + }, "nickname": { "anyOf": [ { @@ -1425,6 +1475,48 @@ "title": "Response[GenerationTaskOut]", "type": "object" }, + "Response_InviteCodeOut_": { + "properties": { + "code": { + "default": 200, + "description": "业务状态码:成功 200,失败非 200", + "title": "Code", + "type": "integer" + }, + "data": { + "anyOf": [ + { + "$ref": "#/components/schemas/InviteCodeOut" + }, + { + "type": "null" + } + ], + "description": "业务数据" + }, + "message": { + "default": "success", + "description": "提示信息", + "title": "Message", + "type": "string" + }, + "timestamp": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "description": "响应时间;默认不携带,不携带时省略", + "title": "Timestamp" + } + }, + "title": "Response[InviteCodeOut]", + "type": "object" + }, "Response_MediaUploadResult_": { "properties": { "code": { @@ -2095,7 +2187,7 @@ }, "/auth/login-by-code": { "post": { - "description": "验证码登录。内测期间不自动注册。", + "description": "验证码登录。未知邮箱自动建号并赠送注册积分。", "operationId": "login_by_code_auth_login_by_code_post", "requestBody": { "content": { @@ -2285,7 +2377,7 @@ }, "/auth/register": { "post": { - "description": "邮箱+验证码+密码注册。\n\n内测期间关闭公开注册,路由与请求模型保留以便以后重新开放。", + "description": "邮箱+验证码+密码注册。邀请码选填。", "operationId": "register_auth_register_post", "requestBody": { "content": { @@ -3106,6 +3198,50 @@ ] } }, + "/quota/invite/code": { + "get": { + "description": "获取当前用户未过期邀请码;没有或已过期则签发新行。", + "operationId": "get_invite_code_quota_invite_code_get", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/Response_InviteCodeOut_" + } + } + }, + "description": "Successful Response" + } + }, + "summary": "Get Invite Code", + "tags": [ + "quota" + ] + } + }, + "/quota/invite/generate": { + "post": { + "description": "签发新邀请码。旧码立即过期,行保留。", + "operationId": "generate_invite_code_quota_invite_generate_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/Response_InviteCodeOut_" + } + } + }, + "description": "Successful Response" + } + }, + "summary": "Generate Invite Code", + "tags": [ + "quota" + ] + } + }, "/quota/transactions": { "get": { "description": "查询积分流水(分页)。",