import uuid from typing import Sequence from sqlalchemy import select, update as sa_update from sqlalchemy.ext.asyncio import AsyncSession from app.models.milestone import Milestone from app.schemas.milestone import MilestoneCreate, MilestoneUpdate async def create(db: AsyncSession, study_id: uuid.UUID, milestone_in: MilestoneCreate) -> Milestone: milestone = Milestone( study_id=study_id, type=milestone_in.type, name=milestone_in.name or milestone_in.type, planned_date=milestone_in.planned_date, actual_date=None, status=milestone_in.status or "NOT_STARTED", owner_id=milestone_in.owner_id, site_id=milestone_in.site_id, notes=milestone_in.notes, ) db.add(milestone) await db.commit() await db.refresh(milestone) return milestone async def get(db: AsyncSession, milestone_id: uuid.UUID) -> Milestone | None: result = await db.execute(select(Milestone).where(Milestone.id == milestone_id)) return result.scalar_one_or_none() async def list_milestones(db: AsyncSession, study_id: uuid.UUID) -> Sequence[Milestone]: result = await db.execute(select(Milestone).where(Milestone.study_id == study_id).order_by(Milestone.planned_date)) return result.scalars().all() async def update(db: AsyncSession, milestone: Milestone, milestone_in: MilestoneUpdate) -> Milestone: update_data = milestone_in.model_dump(exclude_unset=True) if update_data: await db.execute( sa_update(Milestone) .where(Milestone.id == milestone.id) .values(**update_data) ) await db.commit() await db.refresh(milestone) return milestone