import uuid from datetime import date, datetime, timezone from decimal import Decimal from typing import Sequence from sqlalchemy import Numeric, func, select, update as sa_update, case from sqlalchemy.ext.asyncio import AsyncSession from app.models.finance import FinanceItem from app.models.site import Site from app.models.subject import Subject from app.schemas.finance import FinanceCreate, FinanceStatusUpdate, FinanceUpdate VALID_TRANSITIONS = { "DRAFT": {"SUBMITTED"}, "SUBMITTED": {"APPROVED", "REJECTED"}, "APPROVED": {"PAID"}, "REJECTED": set(), "PAID": set(), } async def _validate_site_subject(db: AsyncSession, study_id: uuid.UUID, site_id: uuid.UUID | None, subject_id: uuid.UUID | None): subj_site = None if site_id: 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") if subject_id: result = await db.execute(select(Subject).where(Subject.id == subject_id)) subj = result.scalar_one_or_none() if not subj or subj.study_id != study_id: raise ValueError("Subject not found in study") subj_site = subj.site_id if site_id and subj_site and site_id != subj_site: raise ValueError("Site and subject mismatch") async def create_item(db: AsyncSession, study_id: uuid.UUID, item_in: FinanceCreate, *, created_by: uuid.UUID) -> FinanceItem: await _validate_site_subject(db, study_id, item_in.site_id, item_in.subject_id) item = FinanceItem( study_id=study_id, site_id=item_in.site_id, subject_id=item_in.subject_id, visit_id=item_in.visit_id, category=item_in.category, title=item_in.title, description=item_in.description, currency=item_in.currency, amount=item_in.amount, occur_date=item_in.occur_date, status="DRAFT", created_by=created_by, ) db.add(item) await db.commit() await db.refresh(item) return item async def get_item(db: AsyncSession, item_id: uuid.UUID) -> FinanceItem | None: result = await db.execute(select(FinanceItem).where(FinanceItem.id == item_id)) return result.scalar_one_or_none() async def list_items( db: AsyncSession, study_id: uuid.UUID, *, status: str | None = None, category: str | None = None, site_id: uuid.UUID | None = None, subject_id: uuid.UUID | None = None, date_from: date | None = None, date_to: date | None = None, skip: int = 0, limit: int = 100, ) -> Sequence[FinanceItem]: stmt = select(FinanceItem).where(FinanceItem.study_id == study_id) if status: stmt = stmt.where(FinanceItem.status == status) if category: stmt = stmt.where(FinanceItem.category == category) if site_id: stmt = stmt.where(FinanceItem.site_id == site_id) if subject_id: stmt = stmt.where(FinanceItem.subject_id == subject_id) if date_from: stmt = stmt.where(FinanceItem.occur_date >= date_from) if date_to: stmt = stmt.where(FinanceItem.occur_date <= date_to) stmt = stmt.offset(skip).limit(limit) result = await db.execute(stmt) return result.scalars().all() async def update_item_draft(db: AsyncSession, item: FinanceItem, item_in: FinanceUpdate) -> FinanceItem: if item.status != "DRAFT": raise ValueError("Only DRAFT items can be edited") update_data = item_in.model_dump(exclude_unset=True) if update_data: await db.execute( sa_update(FinanceItem) .where(FinanceItem.id == item.id) .values(**update_data) ) await db.commit() await db.refresh(item) return item async def change_status(db: AsyncSession, item: FinanceItem, status_in: FinanceStatusUpdate, *, operator_id: uuid.UUID) -> FinanceItem: target = status_in.status allowed = VALID_TRANSITIONS.get(item.status, set()) if target not in allowed: raise ValueError(f"Invalid status transition {item.status} -> {target}") update_data = {"status": target} now = datetime.now(timezone.utc) if target == "SUBMITTED": update_data["submitted_at"] = now if target == "APPROVED": update_data["approved_at"] = now update_data["approver_id"] = operator_id update_data["reject_reason"] = None if target == "REJECTED": if not status_in.reject_reason: raise ValueError("reject_reason required") update_data["rejected_at"] = now update_data["approver_id"] = operator_id update_data["reject_reason"] = status_in.reject_reason if target == "PAID": if item.status != "APPROVED": raise ValueError("Only APPROVED can transition to PAID") update_data["paid_at"] = now update_data["payer_id"] = operator_id await db.execute( sa_update(FinanceItem) .where(FinanceItem.id == item.id) .values(**update_data) ) await db.commit() await db.refresh(item) return item async def summary( db: AsyncSession, study_id: uuid.UUID, *, date_from: date | None = None, date_to: date | None = None, ): stmt = select( func.coalesce(func.sum(FinanceItem.amount), 0).label("total_amount"), func.coalesce(func.sum(case((FinanceItem.status == "PAID", FinanceItem.amount), else_=0)), 0).label( "paid_amount" ), func.coalesce(func.sum(case((FinanceItem.status == "APPROVED", FinanceItem.amount), else_=0)), 0).label( "approved_amount" ), func.coalesce(func.sum(case((FinanceItem.status == "SUBMITTED", FinanceItem.amount), else_=0)), 0).label( "submitted_amount" ), func.count().label("count_total"), func.coalesce(func.sum(case((FinanceItem.status == "PAID", 1), else_=0)), 0).label("count_paid"), ).where(FinanceItem.study_id == study_id) if date_from: stmt = stmt.where(FinanceItem.occur_date >= date_from) if date_to: stmt = stmt.where(FinanceItem.occur_date <= date_to) result = await db.execute(stmt) return result.one()