import uuid from datetime import date, timedelta from typing import Sequence from sqlalchemy import select, update as sa_update from sqlalchemy.ext.asyncio import AsyncSession from app.crud import visit as visit_crud from app.models.site import Site from app.models.subject import Subject from app.schemas.subject import SubjectCreate, SubjectUpdate async def _validate_site(db: AsyncSession, study_id: uuid.UUID, site_id: uuid.UUID) -> None: result = await db.execute(select(Site).where(Site.id == site_id)) site = result.scalar_one_or_none() if not site or site.study_id != study_id: raise ValueError("Site not found in study") async def create_subject(db: AsyncSession, study_id: uuid.UUID, subject_in: SubjectCreate) -> Subject: await _validate_site(db, study_id, subject_in.site_id) subject = Subject( study_id=study_id, site_id=subject_in.site_id, subject_no=subject_in.subject_no, status="SCREENING", screening_date=subject_in.screening_date, enrollment_date=None, completion_date=None, drop_reason=None, ) db.add(subject) await db.commit() await db.refresh(subject) # initial visit: Screening (V0) await visit_crud.create_visit( db, study_id=study_id, visit_in=None, subject=subject, visit_code="V0", visit_name="Screening", planned_date=subject.screening_date, ) return subject async def get_subject(db: AsyncSession, subject_id: uuid.UUID) -> Subject | None: result = await db.execute(select(Subject).where(Subject.id == subject_id)) return result.scalar_one_or_none() async def list_subjects( db: AsyncSession, study_id: uuid.UUID, site_id: uuid.UUID | None = None, status: str | None = None, ) -> Sequence[Subject]: stmt = select(Subject).where(Subject.study_id == study_id) if site_id: stmt = stmt.where(Subject.site_id == site_id) if status: stmt = stmt.where(Subject.status == status) result = await db.execute(stmt) return result.scalars().all() async def generate_default_visits(db: AsyncSession, subject: Subject) -> None: # Baseline + Follow-up visits based on enrollment_date if not subject.enrollment_date: return baseline_date = subject.enrollment_date follow1 = baseline_date + timedelta(days=30) follow2 = baseline_date + timedelta(days=60) visits_data = [ ("V1", "Baseline", baseline_date), ("FU1", "Follow-up 1", follow1), ("FU2", "Follow-up 2", follow2), ] for code, name, plan_date in visits_data: await visit_crud.create_visit( db, study_id=subject.study_id, visit_in=None, subject=subject, visit_code=code, visit_name=name, planned_date=plan_date, ) async def update_subject(db: AsyncSession, subject: Subject, subject_in: SubjectUpdate) -> Subject: update_data = subject_in.model_dump(exclude_unset=True) if update_data: await db.execute( sa_update(Subject) .where(Subject.id == subject.id) .values(**update_data) ) await db.commit() await db.refresh(subject) return subject