import uuid from datetime import date from typing import Sequence from sqlalchemy import case, func, or_, select from sqlalchemy.ext.asyncio import AsyncSession from app.models.contract_fee import ContractFee from app.models.contract_fee_payment import ContractFeePayment from app.models.site import Site from app.schemas.contract_fee import ContractFeeCreate, ContractFeeUpdate async def create_contract_fee( db: AsyncSession, contract_in: ContractFeeCreate, ) -> ContractFee: contract = ContractFee( project_id=contract_in.project_id, center_id=contract_in.center_id, contract_no=contract_in.contract_no, signed_date=contract_in.signed_date, contract_amount=contract_in.contract_amount, currency=contract_in.currency, remark=contract_in.remark, contract_cases=contract_in.contract_cases, actual_cases=contract_in.actual_cases, settlement_amount=contract_in.settlement_amount, final_payment_amount=contract_in.final_payment_amount, ) db.add(contract) await db.commit() await db.refresh(contract) return contract async def get_contract_fee(db: AsyncSession, contract_id: uuid.UUID) -> ContractFee | None: result = await db.execute(select(ContractFee).where(ContractFee.id == contract_id)) return result.scalar_one_or_none() async def get_contract_fee_by_project_center( db: AsyncSession, project_id: uuid.UUID, center_id: uuid.UUID ) -> ContractFee | None: result = await db.execute( select(ContractFee).where(ContractFee.project_id == project_id, ContractFee.center_id == center_id) ) return result.scalar_one_or_none() async def list_contract_fees( db: AsyncSession, project_id: uuid.UUID, center_id: uuid.UUID | None = None, center_ids: set[uuid.UUID] | None = None, q: str | None = None, ) -> Sequence[tuple[ContractFee, str, float, float, date | None, date | None]]: paid_total = func.coalesce( func.sum(case((ContractFeePayment.is_paid.is_(True), ContractFeePayment.amount), else_=0)), 0, ).label("paid_total") verified_total = func.coalesce( func.sum(case((ContractFeePayment.is_verified.is_(True), ContractFeePayment.amount), else_=0)), 0, ).label("verified_total") last_paid_date = func.max( case((ContractFeePayment.is_paid.is_(True), ContractFeePayment.paid_date), else_=None) ).label("last_paid_date") last_verified_date = func.max( case((ContractFeePayment.is_verified.is_(True), ContractFeePayment.verified_date), else_=None) ).label("last_verified_date") stmt = ( select( ContractFee, Site.name.label("center_name"), paid_total, verified_total, last_paid_date, last_verified_date, ) .join(Site, Site.id == ContractFee.center_id) .outerjoin(ContractFeePayment, ContractFeePayment.contract_fee_id == ContractFee.id) .where(ContractFee.project_id == project_id) .group_by(ContractFee.id, Site.name) ) if center_id: stmt = stmt.where(ContractFee.center_id == center_id) if center_ids is not None: if not center_ids: return [] stmt = stmt.where(ContractFee.center_id.in_(center_ids)) if q: conditions = [Site.name.ilike(f"%{q}%")] try: conditions.append(Site.id == uuid.UUID(q)) except (ValueError, TypeError): pass stmt = stmt.where(or_(*conditions)) stmt = stmt.order_by(Site.name.asc()) result = await db.execute(stmt) return result.all() async def update_contract_fee( db: AsyncSession, contract: ContractFee, contract_in: ContractFeeUpdate ) -> ContractFee: update_data = contract_in.model_dump(exclude_unset=True) for key, value in update_data.items(): setattr(contract, key, value) await db.commit() await db.refresh(contract) return contract async def delete_contract_fee(db: AsyncSession, contract: ContractFee) -> None: await db.delete(contract) await db.commit()