from __future__ import annotations import uuid from datetime import UTC, datetime from sqlalchemy import select from sqlalchemy.orm import Session from app.models.user import AuditLog, RefreshToken, Role, User, user_roles_table from app.repositories.base import BaseRepository class UserRepository(BaseRepository[User]): """Repository for User operations.""" def __init__(self, db: Session) -> None: super().__init__(db, User) def get_by_username(self, username: str) -> User | None: """Get user by username.""" query = select(User).where(User.username == username) result = self.db.execute(query) return result.scalars().first() def get_by_email(self, email: str) -> User | None: """Get user by email.""" query = select(User).where(User.email == email) result = self.db.execute(query) return result.scalars().first() def create_user( self, username: str, email: str, hashed_password: str, full_name: str | None = None, is_active: bool = True, is_superuser: bool = False, role_names: list[str] | None = None, ) -> User: """Create a new user with optional roles.""" user = User( username=username, email=email, hashed_password=hashed_password, full_name=full_name, is_active=is_active, is_superuser=is_superuser, ) if role_names: roles = self.get_roles_by_names(role_names) user.roles = roles return self.create(user) def update_last_login(self, user: User) -> User: """Update user's last login timestamp.""" user.last_login = datetime.now(UTC) self.db.flush() self.db.refresh(user) return user def get_roles_by_names(self, role_names: list[str]) -> list[Role]: """Get roles by their names.""" query = select(Role).where(Role.name.in_(role_names)) result = self.db.execute(query) return list(result.scalars().all()) def assign_roles(self, user: User, role_names: list[str]) -> User: """Assign roles to a user, replacing existing roles.""" roles = self.get_roles_by_names(role_names) user.roles = roles self.db.flush() self.db.refresh(user) return user def get_active_users(self, offset: int = 0, limit: int = 100) -> list[User]: """Get all active users.""" query = select(User).where(User.is_active.is_(True)).offset(offset).limit(limit) result = self.db.execute(query) return list(result.scalars().all()) class RoleRepository(BaseRepository[Role]): """Repository for Role operations.""" def __init__(self, db: Session) -> None: super().__init__(db, Role) def get_by_name(self, name: str) -> Role | None: """Get role by name.""" query = select(Role).where(Role.name == name) result = self.db.execute(query) return result.scalars().first() def create_role(self, name: str, description: str | None = None) -> Role: """Create a new role.""" role = Role(name=name, description=description) return self.create(role) def get_all_roles(self) -> list[Role]: """Get all roles.""" query = select(Role).order_by(Role.name) result = self.db.execute(query) return list(result.scalars().all()) class RefreshTokenRepository(BaseRepository[RefreshToken]): """Repository for RefreshToken operations.""" def __init__(self, db: Session) -> None: super().__init__(db, RefreshToken) def get_by_token(self, token: str) -> RefreshToken | None: """Get refresh token by token string.""" query = select(RefreshToken).where( RefreshToken.token == token, RefreshToken.revoked.is_(False), RefreshToken.expires_at > datetime.now(UTC), ) result = self.db.execute(query) return result.scalars().first() def create_token(self, user_id: uuid.UUID, token: str, expires_at: datetime) -> RefreshToken: """Create a new refresh token.""" refresh_token = RefreshToken( user_id=user_id, token=token, expires_at=expires_at, ) return self.create(refresh_token) def revoke_token(self, token: str) -> bool: """Revoke a refresh token.""" refresh_token = self.get_by_token(token) if refresh_token: refresh_token.revoked = True self.db.flush() return True return False def revoke_all_user_tokens(self, user_id: uuid.UUID) -> int: """Revoke all refresh tokens for a user.""" query = select(RefreshToken).where( RefreshToken.user_id == user_id, RefreshToken.revoked.is_(False), ) result = self.db.execute(query) tokens = result.scalars().all() count = 0 for token in tokens: token.revoked = True count += 1 self.db.flush() return count def cleanup_expired_tokens(self) -> int: """Remove expired or revoked tokens.""" query = select(RefreshToken).where( (RefreshToken.expires_at <= datetime.now(UTC)) | (RefreshToken.revoked.is_(True)) ) result = self.db.execute(query) tokens = result.scalars().all() count = len(tokens) for token in tokens: self.db.delete(token) self.db.flush() return count class AuditLogRepository(BaseRepository[AuditLog]): """Repository for AuditLog operations.""" def __init__(self, db: Session) -> None: super().__init__(db, AuditLog) def log_action( self, action: str, resource_type: str, resource_id: str | None = None, user_id: uuid.UUID | None = None, details: str | None = None, ip_address: str | None = None, user_agent: str | None = None, ) -> AuditLog: """Create an audit log entry.""" audit_log = AuditLog( user_id=user_id, action=action, resource_type=resource_type, resource_id=resource_id, details=details, ip_address=ip_address, user_agent=user_agent, ) return self.create(audit_log) def get_user_logs( self, user_id: uuid.UUID, offset: int = 0, limit: int = 100, ) -> list[AuditLog]: """Get audit logs for a specific user.""" query = ( select(AuditLog) .where(AuditLog.user_id == user_id) .order_by(AuditLog.created_at.desc()) .offset(offset) .limit(limit) ) result = self.db.execute(query) return list(result.scalars().all()) def get_resource_logs( self, resource_type: str, resource_id: str, offset: int = 0, limit: int = 100, ) -> list[AuditLog]: """Get audit logs for a specific resource.""" query = ( select(AuditLog) .where( AuditLog.resource_type == resource_type, AuditLog.resource_id == resource_id, ) .order_by(AuditLog.created_at.desc()) .offset(offset) .limit(limit) ) result = self.db.execute(query) return list(result.scalars().all())