from typing import Optional from sqlalchemy import delete, func, select, text, update from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.orm import joinedload, selectinload from src.db.models.org_unit import OrgUnit from src.db.models.role import Role from src.db.models.user_org import UserOrg from src.db.models.app_user import AppUser from src.core.security import get_password_hash, verify_password class UserRepository: def __init__(self, db: AsyncSession): self.db = db async def get(self, user_id: int, load_orgs: bool = False) -> Optional[AppUser]: query = select(AppUser).where(AppUser.id == user_id) if load_orgs: query = query.options( selectinload(AppUser.org_units), ) return ( ( await self.db.execute( query.limit(1) ) ) .scalars() .first() ) async def get_by_email(self, email: str) -> Optional[AppUser]: return ( ( await self.db.execute( select(AppUser).where(AppUser.email == email).limit(1) ) ) .scalars() .first() ) async def get_by_username(self, username: str) -> Optional[AppUser]: return ( ( await self.db.execute( select(AppUser).where(AppUser.username == username).limit(1) ) ) .scalars() .first() ) async def set_app_user_id(self, user_id: int) -> None: if isinstance(user_id, int): await self.db.execute( text( f"SET LOCAL app.user_id = {user_id}" ), ) async def get_list( self, skip: int = 0, limit: int = 100, with_count: bool = False, load_orgs: bool = False, ) -> list[AppUser] | tuple[int, list[AppUser]]: where = [AppUser.is_active.is_(True)] if with_count: query = select(func.count().over().label("total_count"), AppUser) else: query = select(AppUser) query = query.where(*where) if load_orgs: query = query.options( joinedload(AppUser.org_units) ) query = query.order_by(AppUser.id) if skip is not None: query = query.offset(skip) if limit is not None: query = query.limit(limit) if with_count: result = (await self.db.execute(query)).unique().all() if result: return result[0][0], [res[1] for res in result] else: count_query = select(func.count(AppUser.id)) if query.whereclause is not None: count_query = count_query.where( query.whereclause ) return (await self.db.execute(count_query)).scalar(), [] else: return (await self.db.execute(query)).scalars().unique().all() async def logical_delete(self, user_id: int) -> bool: query = update(AppUser).where(AppUser.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[AppUser]: 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[AppUser]: 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[OrgUnit]: query = select(OrgUnit).join(UserOrg, UserOrg.org_unit_id == OrgUnit.id).where( UserOrg.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(UserOrg.org_unit_id).where(UserOrg.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: for ssp_id in ssp_ids: await self.db.execute( text( "SELECT v3.grant_user_org_access(:user_id, :org_unit_id)" ), {"user_id": user_id, "org_unit_id": ssp_id}, ) await self.db.commit() return True async def unset_many_ssp(self, user_id: int, ssp_ids: list[int]) -> bool: for ssp_id in ssp_ids: await self.db.execute( text( "SELECT v3.revoke_user_org_access(:user_id, :org_unit_id)" ), {"user_id": user_id, "org_unit_id": ssp_id}, ) await self.db.commit() return True async def clear_many_ssp(self, user_id: int) -> None: org_unit_ids = await self.get_many_ssp_ids(user_id) await self.db.execute( text( "SELECT v3.revoke_many_user_org_access(:user_id, :org_unit_ids)" ), {"user_id": user_id, "org_unit_ids": org_unit_ids}, ) await self.db.flush() async def create(self, user_data: dict) -> AppUser: hashed_password = ( get_password_hash(user_data["password"]) if user_data.get("password") else None ) user_id = ( await self.db.execute( text( """ SELECT (v3.add_user( :email, :username, :hashed_password, :full_name, :role_id )).id """ ), { "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"), }, ) ).scalar_one() await self.db.flush() return await self.get(user_id) async def update( self, user_id: int, user_data: dict, load_orgs: bool = False, ) -> Optional[AppUser]: user = await self.get(user_id) if not user: return None hashed_password = None if "password" in user_data and user_data["password"]: hashed_password = get_password_hash(user_data["password"]) await self.db.execute( text( """ SELECT (v3.upd_user( :user_id, :email, :username, :hashed_password, :full_name, :role_id, :is_active )).id """ ), { "user_id": user_id, "email": user_data.get("email"), "username": user_data.get("username"), "hashed_password": hashed_password, "full_name": user_data.get("full_name"), "role_id": user_data.get("role_id"), "is_active": user_data.get("is_active"), }, ) await self.db.flush() self.db.expire(user) result = await self.get(user_id, load_orgs=load_orgs) return result