"""User 用户 Repository。""" from sqlalchemy import select, or_ from sqlalchemy.orm import Session from app.models.user import User class UserRepository: def __init__(self, session: Session) -> None: self._session = session def get_by_id(self, user_id: str) -> User | None: return self._session.get(User, user_id) def get_by_username(self, username: str) -> User | None: stmt = select(User).where(User.username == username) return self._session.scalars(stmt).first() def get_by_email(self, email: str) -> User | None: stmt = select(User).where(User.email == email) return self._session.scalars(stmt).first() def get_by_username_or_email(self, value: str) -> User | None: """登录用:按用户名或邮箱查找。""" stmt = select(User).where( or_(User.username == value, User.email == value.lower()) ) return self._session.scalars(stmt).first() def exists_username_or_email(self, username: str, email: str) -> tuple[bool, bool]: """检查用户名/邮箱是否已存在。返回 (username_taken, email_taken)。""" stmt = select(User).where( or_(User.username == username, User.email == email.lower()) ) existing = list(self._session.scalars(stmt).all()) return ( any(u.username == username for u in existing), any(u.email == email.lower() for u in existing), ) def create( self, *, username: str, email: str, password_hash: str, plan_id: str, ) -> User: user = User( username=username, email=email.lower(), password_hash=password_hash, status="active", plan_id=plan_id, storage_used=0, ) self._session.add(user) self._session.flush() return user def update_password(self, user: User, password_hash: str) -> None: user.password_hash = password_hash self._session.flush()