66 lines
2.1 KiB
Python
66 lines
2.1 KiB
Python
"""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,
|
|
role: str = "customer",
|
|
) -> User:
|
|
user = User(
|
|
username=username,
|
|
email=email.lower(),
|
|
password_hash=password_hash,
|
|
status="active",
|
|
role=role,
|
|
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() |