237 lines
7.4 KiB
Python
237 lines
7.4 KiB
Python
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())
|