finance_budget.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298
  1. """Budget validation helpers for finance-aware print dispatch."""
  2. import calendar
  3. from datetime import datetime, timezone
  4. from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
  5. from fastapi import HTTPException
  6. from sqlalchemy import case, func, select
  7. from sqlalchemy.ext.asyncio import AsyncSession
  8. from backend.app.models.finance import BudgetReservation, CostCenter, CostCenterMember, WalletTransaction
  9. from backend.app.models.print_queue import PrintQueueItem
  10. from backend.app.models.settings import Settings
  11. from backend.app.models.user import User
  12. async def is_billing_enabled(db: AsyncSession) -> bool:
  13. # Consider any 'billing_enabled' setting with a true-ish value as enabling billing.
  14. result = await db.execute(
  15. select(func.count())
  16. .select_from(Settings)
  17. .where(Settings.key == "billing_enabled", func.lower(func.coalesce(Settings.value, "")) == "true")
  18. )
  19. count = int(result.scalar_one() or 0)
  20. return count > 0
  21. async def is_printer_kill_switch_enabled(db: AsyncSession) -> bool:
  22. """Return True when billing and the printer kill-switch are both enabled."""
  23. result = await db.execute(
  24. select(Settings.key, Settings.value).where(Settings.key.in_(("billing_enabled", "printer_kill_switch_enabled")))
  25. )
  26. values = {key: (value or "").strip().lower() for key, value in result.all()}
  27. return values.get("billing_enabled") == "true" and values.get("printer_kill_switch_enabled") == "true"
  28. async def _get_budget_window_start_utc(db: AsyncSession) -> datetime:
  29. result = await db.execute(
  30. select(Settings).where(Settings.key.in_(["finance_budget_reset_day", "finance_budget_reset_timezone"]))
  31. )
  32. values = {setting.key: setting.value for setting in result.scalars().all()}
  33. desired_day = 1
  34. try:
  35. parsed = int(values.get("finance_budget_reset_day") or 1)
  36. if 1 <= parsed <= 31:
  37. desired_day = parsed
  38. except (TypeError, ValueError):
  39. pass
  40. timezone_name = values.get("finance_budget_reset_timezone") or "UTC"
  41. try:
  42. tz = ZoneInfo(timezone_name)
  43. except ZoneInfoNotFoundError:
  44. tz = ZoneInfo("UTC")
  45. now = datetime.now(tz)
  46. current_month_reset_day = min(desired_day, calendar.monthrange(now.year, now.month)[1])
  47. if now.day < current_month_reset_day:
  48. month = now.month - 1
  49. year = now.year
  50. if month == 0:
  51. month = 12
  52. year -= 1
  53. else:
  54. month = now.month
  55. year = now.year
  56. reset_day = min(desired_day, calendar.monthrange(year, month)[1])
  57. return datetime(year, month, reset_day, tzinfo=tz).astimezone(timezone.utc)
  58. async def _cost_center_spend(db: AsyncSession, cost_center_id: int, *, monthly: bool) -> float:
  59. spend_expr = case((WalletTransaction.amount < 0, -WalletTransaction.amount), else_=0.0)
  60. conditions = [
  61. WalletTransaction.cost_center_id == cost_center_id,
  62. WalletTransaction.cost_center_id.is_not(None),
  63. WalletTransaction.is_voided.is_(False),
  64. ]
  65. if monthly:
  66. conditions.append(WalletTransaction.created_at >= await _get_budget_window_start_utc(db))
  67. result = await db.execute(select(func.coalesce(func.sum(spend_expr), 0.0)).where(*conditions))
  68. return float(result.scalar() or 0.0)
  69. async def get_cost_center_reserved_map(
  70. db: AsyncSession,
  71. cost_center_ids: list[int],
  72. *,
  73. exclude_queue_item_id: int | None = None,
  74. exclude_reservation_source_type: str | None = None,
  75. exclude_reservation_source_id: int | None = None,
  76. ) -> dict[int, float]:
  77. """Return active holds plus unreserved open queue estimates per cost center.
  78. Queue items that already have an active ``print_queue`` reservation are
  79. excluded from the queue sum because the reservation is their replacement,
  80. not an additional hold.
  81. """
  82. if not cost_center_ids:
  83. return {}
  84. active_queue_reservation = (
  85. select(BudgetReservation.id)
  86. .where(
  87. BudgetReservation.status == "active",
  88. BudgetReservation.source_type == "print_queue",
  89. BudgetReservation.source_id == PrintQueueItem.id,
  90. )
  91. .exists()
  92. )
  93. queue_conditions = [
  94. PrintQueueItem.cost_center_id.in_(cost_center_ids),
  95. PrintQueueItem.status.in_(("pending", "printing")),
  96. ~active_queue_reservation,
  97. ]
  98. if exclude_queue_item_id is not None:
  99. queue_conditions.append(PrintQueueItem.id != exclude_queue_item_id)
  100. queue_rows = await db.execute(
  101. select(PrintQueueItem.cost_center_id, func.coalesce(func.sum(PrintQueueItem.estimated_cost), 0.0))
  102. .where(*queue_conditions)
  103. .group_by(PrintQueueItem.cost_center_id)
  104. )
  105. reserved_map = {int(center_id): float(value) for center_id, value in queue_rows.all() if center_id is not None}
  106. reservation_conditions = [
  107. BudgetReservation.cost_center_id.in_(cost_center_ids),
  108. BudgetReservation.status == "active",
  109. ]
  110. if exclude_reservation_source_type is not None and exclude_reservation_source_id is not None:
  111. reservation_conditions.append(
  112. ~(
  113. (BudgetReservation.source_type == exclude_reservation_source_type)
  114. & (BudgetReservation.source_id == exclude_reservation_source_id)
  115. )
  116. )
  117. reservation_rows = await db.execute(
  118. select(BudgetReservation.cost_center_id, func.coalesce(func.sum(BudgetReservation.amount), 0.0))
  119. .where(*reservation_conditions)
  120. .group_by(BudgetReservation.cost_center_id)
  121. )
  122. for center_id, value in reservation_rows.all():
  123. if center_id is not None:
  124. reserved_map[int(center_id)] = reserved_map.get(int(center_id), 0.0) + float(value or 0.0)
  125. return reserved_map
  126. async def validate_print_budget(
  127. db: AsyncSession,
  128. *,
  129. cost_center_id: int | None,
  130. estimated_cost: float | None,
  131. current_user: User | None,
  132. quantity: int = 1,
  133. exclude_queue_item_id: int | None = None,
  134. exclude_reservation_source_type: str | None = None,
  135. exclude_reservation_source_id: int | None = None,
  136. ) -> None:
  137. """Validate that a print can be assigned to a cost center budget."""
  138. if not await is_billing_enabled(db):
  139. return
  140. if cost_center_id is None:
  141. raise HTTPException(status_code=400, detail="Cost center is required when billing is enabled")
  142. if estimated_cost is None or estimated_cost <= 0:
  143. raise HTTPException(status_code=400, detail="Estimated cost is required for cost center prints")
  144. center = await db.scalar(select(CostCenter).where(CostCenter.id == cost_center_id).with_for_update())
  145. if not center:
  146. raise HTTPException(status_code=404, detail="Cost center not found")
  147. if not center.is_active:
  148. raise HTTPException(status_code=400, detail="Cost center is inactive")
  149. if current_user is not None and not current_user.is_admin:
  150. if center.is_private:
  151. if center.owner_user_id != current_user.id:
  152. raise HTTPException(status_code=403, detail="You cannot print with this private cost center")
  153. else:
  154. member = await db.scalar(
  155. select(CostCenterMember).where(
  156. CostCenterMember.cost_center_id == cost_center_id,
  157. CostCenterMember.user_id == current_user.id,
  158. )
  159. )
  160. if not member or not member.can_print:
  161. raise HTTPException(status_code=403, detail="You cannot print with this cost center")
  162. budget_limit = center.monthly_budget if center.monthly_budget is not None else center.total_budget
  163. if budget_limit is None:
  164. return
  165. used = await _cost_center_spend(db, cost_center_id, monthly=center.monthly_budget is not None)
  166. reserved_map = await get_cost_center_reserved_map(
  167. db,
  168. [cost_center_id],
  169. exclude_queue_item_id=exclude_queue_item_id,
  170. exclude_reservation_source_type=exclude_reservation_source_type,
  171. exclude_reservation_source_id=exclude_reservation_source_id,
  172. )
  173. reserved = reserved_map.get(cost_center_id, 0.0)
  174. requested = estimated_cost * max(1, quantity)
  175. available = float(budget_limit) - used - reserved
  176. if requested > available:
  177. raise HTTPException(
  178. status_code=400,
  179. detail=f"Estimated print cost exceeds available cost center budget ({requested:.2f} > {available:.2f})",
  180. )
  181. async def create_budget_reservation(
  182. db: AsyncSession,
  183. *,
  184. cost_center_id: int | None,
  185. estimated_cost: float | None,
  186. current_user: User | None,
  187. source_type: str,
  188. source_id: int | None,
  189. print_archive_id: int | None = None,
  190. exclude_queue_item_id: int | None = None,
  191. ) -> BudgetReservation | None:
  192. if not await is_billing_enabled(db):
  193. return None
  194. if cost_center_id is None:
  195. raise HTTPException(status_code=400, detail="Cost center is required when billing is enabled")
  196. await validate_print_budget(
  197. db,
  198. cost_center_id=cost_center_id,
  199. estimated_cost=estimated_cost,
  200. current_user=current_user,
  201. exclude_queue_item_id=exclude_queue_item_id,
  202. exclude_reservation_source_type=source_type,
  203. exclude_reservation_source_id=source_id,
  204. )
  205. existing = None
  206. if source_id is not None:
  207. existing = await db.scalar(
  208. select(BudgetReservation).where(
  209. BudgetReservation.status == "active",
  210. BudgetReservation.source_type == source_type,
  211. BudgetReservation.source_id == source_id,
  212. )
  213. )
  214. if existing is not None:
  215. existing.cost_center_id = cost_center_id
  216. existing.amount = float(estimated_cost or 0.0)
  217. if print_archive_id is not None:
  218. existing.print_archive_id = print_archive_id
  219. await db.flush()
  220. return existing
  221. reservation = BudgetReservation(
  222. cost_center_id=cost_center_id,
  223. amount=float(estimated_cost or 0.0),
  224. status="active",
  225. source_type=source_type,
  226. source_id=source_id,
  227. print_archive_id=print_archive_id,
  228. )
  229. db.add(reservation)
  230. await db.flush()
  231. return reservation
  232. async def release_budget_reservation(
  233. db: AsyncSession,
  234. *,
  235. source_type: str | None = None,
  236. source_id: int | None = None,
  237. print_archive_id: int | None = None,
  238. status: str = "released",
  239. ) -> int:
  240. conditions = [BudgetReservation.status == "active"]
  241. if print_archive_id is not None:
  242. conditions.append(BudgetReservation.print_archive_id == print_archive_id)
  243. else:
  244. conditions.extend(
  245. [
  246. BudgetReservation.source_type == source_type,
  247. BudgetReservation.source_id == source_id,
  248. ]
  249. )
  250. result = await db.execute(select(BudgetReservation).where(*conditions))
  251. reservations = result.scalars().all()
  252. for reservation in reservations:
  253. reservation.status = status
  254. reservation.released_at = datetime.now(timezone.utc)
  255. if reservations:
  256. await db.flush()
  257. return len(reservations)