175 lines
6.1 KiB
Python
175 lines
6.1 KiB
Python
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()
|