Files
OCR/docengine/app/repositories/user_repository.py
2026-06-01 21:49:53 +05:30

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())