98 lines
2.8 KiB
Python
98 lines
2.8 KiB
Python
from sqlalchemy import select
|
|
from sqlalchemy.orm import Session
|
|
|
|
from app.core.security import hash_password
|
|
from app.models.user import User
|
|
from app.schemas.user import UserCreate, UserPasswordUpdate, UserUpdate
|
|
|
|
|
|
class UserRepository:
|
|
@staticmethod
|
|
def get_all(db: Session) -> list[User]:
|
|
return list(db.scalars(select(User).order_by(User.created_at.desc())))
|
|
|
|
@staticmethod
|
|
def get_by_id(db: Session, user_id: int) -> User | None:
|
|
return db.get(User, user_id)
|
|
|
|
@staticmethod
|
|
def get_by_username(db: Session, username: str) -> User | None:
|
|
return db.scalar(select(User).where(User.username == username))
|
|
|
|
@staticmethod
|
|
def find_conflict(
|
|
db: Session,
|
|
*,
|
|
username: str,
|
|
email: str,
|
|
exclude_user_id: int | None = None,
|
|
) -> tuple[str, User] | None:
|
|
username_query = select(User).where(User.username == username)
|
|
email_query = select(User).where(User.email == email)
|
|
|
|
if exclude_user_id is not None:
|
|
username_query = username_query.where(User.id != exclude_user_id)
|
|
email_query = email_query.where(User.id != exclude_user_id)
|
|
|
|
username_user = db.scalar(username_query)
|
|
if username_user is not None:
|
|
return ("username", username_user)
|
|
|
|
email_user = db.scalar(email_query)
|
|
if email_user is not None:
|
|
return ("email", email_user)
|
|
|
|
return None
|
|
|
|
@staticmethod
|
|
def create(db: Session, user: UserCreate) -> User:
|
|
db_user = User(
|
|
first_name=user.first_name,
|
|
last_name=user.last_name,
|
|
username=user.username,
|
|
email=str(user.email),
|
|
role=user.role,
|
|
is_active=user.is_active,
|
|
password_hash=hash_password(user.password),
|
|
)
|
|
|
|
db.add(db_user)
|
|
db.commit()
|
|
db.refresh(db_user)
|
|
|
|
return db_user
|
|
|
|
@staticmethod
|
|
def update(db: Session, db_user: User, user: UserUpdate) -> User:
|
|
db_user.first_name = user.first_name
|
|
db_user.last_name = user.last_name
|
|
db_user.username = user.username
|
|
db_user.email = str(user.email)
|
|
db_user.role = user.role
|
|
db_user.is_active = user.is_active
|
|
|
|
if user.password:
|
|
db_user.password_hash = hash_password(user.password)
|
|
|
|
db.commit()
|
|
db.refresh(db_user)
|
|
|
|
return db_user
|
|
|
|
@staticmethod
|
|
def update_password(
|
|
db: Session,
|
|
db_user: User,
|
|
password_update: UserPasswordUpdate,
|
|
) -> User:
|
|
db_user.password_hash = hash_password(password_update.password)
|
|
|
|
db.commit()
|
|
db.refresh(db_user)
|
|
|
|
return db_user
|
|
|
|
@staticmethod
|
|
def delete(db: Session, db_user: User) -> None:
|
|
db.delete(db_user)
|
|
db.commit()
|