150 lines
4.9 KiB
Python
150 lines
4.9 KiB
Python
from typing import Optional
|
|
|
|
from sqlalchemy import delete, select, update
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from src.core.security import get_password_hash, verify_password
|
|
from src.domain.models import Role, SSP, UserSSPLink, Users
|
|
|
|
|
|
class UserRepository:
|
|
def __init__(self, db: AsyncSession):
|
|
self.db = db
|
|
|
|
async def get(self, user_id: int) -> Optional[Users]:
|
|
return (
|
|
(
|
|
await self.db.execute(
|
|
select(Users).where(Users.id == user_id).limit(1)
|
|
)
|
|
)
|
|
.scalars()
|
|
.first()
|
|
)
|
|
|
|
async def get_by_email(self, email: str) -> Optional[Users]:
|
|
return (
|
|
(
|
|
await self.db.execute(
|
|
select(Users).where(Users.email == email).limit(1)
|
|
)
|
|
)
|
|
.scalars()
|
|
.first()
|
|
)
|
|
|
|
async def get_by_username(self, username: str) -> Optional[Users]:
|
|
return (
|
|
(
|
|
await self.db.execute(
|
|
select(Users).where(Users.username == username).limit(1)
|
|
)
|
|
)
|
|
.scalars()
|
|
.first()
|
|
)
|
|
|
|
async def get_list(self, skip: int = 0, limit: int = 100) -> list[Users]:
|
|
query = (
|
|
select(Users)
|
|
.where(Users.is_active.is_(True))
|
|
.offset(skip)
|
|
.limit(limit)
|
|
.order_by(Users.id)
|
|
)
|
|
return (await self.db.execute(query)).scalars().all()
|
|
|
|
async def create(self, user_data: dict) -> Users:
|
|
hashed_password = (
|
|
get_password_hash(user_data["password"]) if user_data.get("password") else None
|
|
)
|
|
db_user = Users(
|
|
email=user_data["email"],
|
|
username=user_data["username"],
|
|
hashed_password=hashed_password,
|
|
full_name=user_data.get("full_name"),
|
|
role_id=user_data.get("role_id"),
|
|
)
|
|
self.db.add(db_user)
|
|
await self.db.commit()
|
|
await self.db.refresh(db_user)
|
|
return db_user
|
|
|
|
async def update(self, user_id: int, user_data: dict) -> Optional[Users]:
|
|
user = await self.get(user_id)
|
|
if not user:
|
|
return None
|
|
|
|
for field, value in user_data.items():
|
|
if field == "password" and value:
|
|
value = get_password_hash(value)
|
|
setattr(user, field, value)
|
|
|
|
await self.db.commit()
|
|
await self.db.refresh(user)
|
|
return user
|
|
|
|
async def logical_delete(self, user_id: int) -> bool:
|
|
query = update(Users).where(Users.id == user_id).values(is_active=False)
|
|
result = await self.db.execute(query)
|
|
await self.db.commit()
|
|
return result.rowcount > 0
|
|
|
|
async def authenticate(self, username: str, password: str) -> Optional[Users]:
|
|
user = await self.get_by_username(username)
|
|
if not user:
|
|
return None
|
|
if not verify_password(password, user.hashed_password):
|
|
return None
|
|
return user
|
|
|
|
async def authenticate_via_email(self, email: str) -> Optional[Users]:
|
|
return await self.get_by_email(email)
|
|
|
|
async def get_roles(self) -> list[Role]:
|
|
return (await self.db.execute(select(Role).order_by(Role.id))).scalars().all()
|
|
|
|
async def get_many_ssp(self, user_id: int) -> list[SSP]:
|
|
query = select(SSP).join(UserSSPLink, UserSSPLink.ssp_id == SSP.id).where(
|
|
UserSSPLink.user_id == user_id
|
|
)
|
|
return (await self.db.execute(query)).scalars().all()
|
|
|
|
async def get_many_ssp_ids(self, user_id: int) -> list[int]:
|
|
query = select(UserSSPLink.ssp_id).where(UserSSPLink.user_id == user_id)
|
|
return (await self.db.execute(query)).scalars().all()
|
|
|
|
async def set_many_ssp(self, user_id: int, ssp_ids: list[int]) -> bool:
|
|
try:
|
|
for ssp_id in ssp_ids:
|
|
link = UserSSPLink(user_id=user_id, ssp_id=ssp_id)
|
|
self.db.add(link)
|
|
await self.db.commit()
|
|
return True
|
|
except Exception:
|
|
await self.db.rollback()
|
|
return False
|
|
|
|
async def unset_many_ssp(self, user_id: int, ssp_ids: list[int]) -> bool:
|
|
try:
|
|
for ssp_id in ssp_ids:
|
|
query = delete(UserSSPLink).where(
|
|
(UserSSPLink.user_id == user_id) & (UserSSPLink.ssp_id == ssp_id)
|
|
)
|
|
await self.db.execute(query)
|
|
await self.db.commit()
|
|
return True
|
|
except Exception:
|
|
await self.db.rollback()
|
|
return False
|
|
|
|
async def clear_many_ssp(self, user_id: int) -> bool:
|
|
try:
|
|
query = delete(UserSSPLink).where(UserSSPLink.user_id == user_id)
|
|
await self.db.execute(query)
|
|
await self.db.commit()
|
|
return True
|
|
except Exception:
|
|
await self.db.rollback()
|
|
return False
|