207 lines
6.7 KiB
Python
207 lines
6.7 KiB
Python
from __future__ import annotations
|
|
|
|
import uuid
|
|
from typing import Sequence
|
|
|
|
from sqlalchemy import delete, func, or_, select, update
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.core.config import (
|
|
PROTECTED_ADMIN_CLINICAL_DEPARTMENT,
|
|
PROTECTED_ADMIN_DEFAULT_PASSWORD,
|
|
PROTECTED_ADMIN_EMAIL,
|
|
PROTECTED_ADMIN_FULL_NAME,
|
|
)
|
|
from app.core.security import hash_password
|
|
from app.models.audit_log import AuditLog
|
|
from app.models.permission_access_log import PermissionAccessLog
|
|
from app.models.study_member import StudyMember
|
|
from app.models.user import User, UserStatus
|
|
from app.schemas.user import UserCreate, UserRegisterRequest, UserUpdate
|
|
|
|
|
|
async def get_by_email(db: AsyncSession, email: str) -> User | None:
|
|
result = await db.execute(select(User).where(User.email == email))
|
|
return result.scalar_one_or_none()
|
|
|
|
|
|
def is_protected_admin_email(email: str | None) -> bool:
|
|
return (email or "").strip().lower() == PROTECTED_ADMIN_EMAIL
|
|
|
|
|
|
def is_protected_admin_user(user: User | None) -> bool:
|
|
return user is not None and is_protected_admin_email(user.email)
|
|
|
|
|
|
async def get_by_id(db: AsyncSession, user_id: uuid.UUID) -> User | None:
|
|
result = await db.execute(select(User).where(User.id == user_id))
|
|
return result.scalar_one_or_none()
|
|
|
|
|
|
async def create_user(
|
|
db: AsyncSession, user_in: UserCreate, *, status: UserStatus | None = UserStatus.ACTIVE
|
|
) -> User:
|
|
status_value = status or UserStatus(user_in.status)
|
|
user = User(
|
|
email=user_in.email,
|
|
password_hash=hash_password(user_in.password),
|
|
full_name=user_in.full_name,
|
|
clinical_department=user_in.clinical_department,
|
|
status=status_value,
|
|
)
|
|
db.add(user)
|
|
await db.commit()
|
|
await db.refresh(user)
|
|
return user
|
|
|
|
|
|
async def create_registered_user(db: AsyncSession, user_in: UserRegisterRequest) -> User:
|
|
return await create_user(db, UserCreate(**user_in.model_dump()), status=UserStatus.ACTIVE)
|
|
|
|
|
|
async def update_user(db: AsyncSession, user: User, user_in: UserUpdate) -> User:
|
|
update_data = {}
|
|
if user_in.status is not None:
|
|
update_data["status"] = UserStatus(user_in.status)
|
|
if user_in.is_active is not None:
|
|
update_data["status"] = UserStatus.ACTIVE if user_in.is_active else UserStatus.DISABLED
|
|
if user_in.avatar_url is not None:
|
|
update_data["avatar_url"] = user_in.avatar_url
|
|
if user_in.password:
|
|
update_data["password_hash"] = hash_password(user_in.password)
|
|
if user_in.full_name is not None:
|
|
update_data["full_name"] = user_in.full_name
|
|
if user_in.clinical_department is not None:
|
|
update_data["clinical_department"] = user_in.clinical_department
|
|
|
|
if update_data:
|
|
await db.execute(update(User).where(User.id == user.id).values(**update_data))
|
|
await db.commit()
|
|
await db.refresh(user)
|
|
return user
|
|
|
|
|
|
def _apply_user_filters(query, *, keyword: str | None = None, status: UserStatus | None = None):
|
|
if keyword:
|
|
pattern = f"%{keyword.strip()}%"
|
|
query = query.where(
|
|
or_(
|
|
User.email.ilike(pattern),
|
|
User.full_name.ilike(pattern),
|
|
User.clinical_department.ilike(pattern),
|
|
)
|
|
)
|
|
if status is not None:
|
|
query = query.where(User.status == status)
|
|
return query
|
|
|
|
|
|
async def list_users(
|
|
db: AsyncSession,
|
|
skip: int = 0,
|
|
limit: int = 100,
|
|
*,
|
|
keyword: str | None = None,
|
|
status: UserStatus | None = None,
|
|
) -> Sequence[User]:
|
|
query = _apply_user_filters(select(User), keyword=keyword, status=status)
|
|
result = await db.execute(query.order_by(User.created_at.desc()).offset(skip).limit(limit))
|
|
return result.scalars().all()
|
|
|
|
|
|
async def count_users(
|
|
db: AsyncSession,
|
|
*,
|
|
keyword: str | None = None,
|
|
status: UserStatus | None = None,
|
|
) -> int:
|
|
query = _apply_user_filters(select(func.count()).select_from(User), keyword=keyword, status=status)
|
|
result = await db.execute(query)
|
|
return int(result.scalar_one() or 0)
|
|
|
|
|
|
async def list_active_member_candidates_for_study(
|
|
db: AsyncSession,
|
|
study_id: uuid.UUID,
|
|
skip: int = 0,
|
|
limit: int = 100,
|
|
) -> Sequence[User]:
|
|
existing_member = (
|
|
select(StudyMember.id)
|
|
.where(StudyMember.study_id == study_id, StudyMember.user_id == User.id)
|
|
.exists()
|
|
)
|
|
stmt = (
|
|
select(User)
|
|
.where(User.status == UserStatus.ACTIVE)
|
|
.where(~existing_member)
|
|
.order_by(User.full_name.asc(), User.email.asc())
|
|
.offset(skip)
|
|
.limit(limit)
|
|
)
|
|
result = await db.execute(stmt)
|
|
return result.scalars().all()
|
|
|
|
|
|
async def count_active_admins(db: AsyncSession) -> int:
|
|
result = await db.execute(
|
|
select(func.count())
|
|
.select_from(User)
|
|
.where(User.is_admin.is_(True), User.status == UserStatus.ACTIVE)
|
|
)
|
|
return int(result.scalar_one() or 0)
|
|
|
|
|
|
async def user_has_retained_history(db: AsyncSession, user_id: uuid.UUID) -> bool:
|
|
audit_count = await db.scalar(
|
|
select(func.count()).select_from(AuditLog).where(AuditLog.operator_id == user_id)
|
|
)
|
|
if int(audit_count or 0) > 0:
|
|
return True
|
|
|
|
permission_log_count = await db.scalar(
|
|
select(func.count()).select_from(PermissionAccessLog).where(PermissionAccessLog.user_id == user_id)
|
|
)
|
|
return int(permission_log_count or 0) > 0
|
|
|
|
|
|
async def list_users_by_status(
|
|
db: AsyncSession, status: UserStatus | None = None, skip: int = 0, limit: int = 100
|
|
) -> Sequence[User]:
|
|
stmt = select(User).offset(skip).limit(limit)
|
|
if status:
|
|
stmt = stmt.where(User.status == status)
|
|
result = await db.execute(stmt.order_by(User.created_at.desc()))
|
|
return result.scalars().all()
|
|
|
|
|
|
async def ensure_admin_exists(db: AsyncSession, *, default_password: str = PROTECTED_ADMIN_DEFAULT_PASSWORD) -> None:
|
|
result = await db.execute(select(User).where(User.email == PROTECTED_ADMIN_EMAIL))
|
|
admin = result.scalar_one_or_none()
|
|
if admin:
|
|
return
|
|
new_admin = User(
|
|
email=PROTECTED_ADMIN_EMAIL,
|
|
password_hash=hash_password(default_password),
|
|
full_name=PROTECTED_ADMIN_FULL_NAME,
|
|
is_admin=True,
|
|
clinical_department=PROTECTED_ADMIN_CLINICAL_DEPARTMENT,
|
|
status=UserStatus.ACTIVE,
|
|
)
|
|
db.add(new_admin)
|
|
await db.commit()
|
|
|
|
|
|
async def get_users_by_ids(db: AsyncSession, ids: set[uuid.UUID]) -> dict[uuid.UUID, User]:
|
|
if not ids:
|
|
return {}
|
|
result = await db.execute(select(User).where(User.id.in_(ids)))
|
|
users = result.scalars().all()
|
|
return {u.id: u for u in users}
|
|
|
|
|
|
async def delete_user(db: AsyncSession, user: User) -> None:
|
|
await db.execute(delete(StudyMember).where(StudyMember.user_id == user.id))
|
|
await db.delete(user)
|
|
await db.commit()
|