sponsor_prompt.py 8.9 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256
  1. """Sponsor-prompt trigger evaluator and dismiss handler.
  2. Drives the in-app "support keeps Bambuddy independent" toast. Trigger families
  3. fire at milestones the user has earned (prints, archives, filament cost,
  4. anniversary) plus a soft version-update nudge after a major upgrade.
  5. A 14-day cooldown applies across ALL families: if any toast fired in the last
  6. 14 days, no new toast fires. Each individual milestone is shown at most once
  7. per user (or once per install in auth-disabled mode); version-update is the
  8. exception — it re-arms every time the running version is newer than the one
  9. last acknowledged.
  10. """
  11. from __future__ import annotations
  12. import json
  13. import logging
  14. from dataclasses import dataclass, field
  15. from datetime import datetime, timedelta, timezone
  16. from typing import Any
  17. from sqlalchemy import func, select
  18. from sqlalchemy.ext.asyncio import AsyncSession
  19. from backend.app.core.config import APP_VERSION
  20. from backend.app.models.archive import PrintArchive
  21. from backend.app.models.print_log import PrintLogEntry
  22. from backend.app.models.sponsor_toast_state import SponsorToastState
  23. from backend.app.models.user import User
  24. logger = logging.getLogger(__name__)
  25. COOLDOWN_DAYS = 14
  26. PRINT_MILESTONES = (10, 25, 100, 500, 1000, 2500, 5000)
  27. COST_MILESTONES = (25, 50, 100, 500, 1000)
  28. ARCHIVE_MILESTONES = (5, 10, 50, 250, 1000)
  29. ANNIVERSARY_YEARS = 1
  30. @dataclass
  31. class Trigger:
  32. """Evaluated trigger result returned to the frontend."""
  33. milestone: str # e.g. "prints-500", "anniversary-1", "version-update"
  34. family: str # "prints" | "cost" | "archives" | "anniversary" | "version-update"
  35. threshold: int | None = None
  36. payload: dict[str, Any] = field(default_factory=dict)
  37. # ---------------------------------------------------------------------------
  38. # State helpers
  39. # ---------------------------------------------------------------------------
  40. async def _get_or_create_state(db: AsyncSession, user_id: int | None) -> SponsorToastState:
  41. """Fetch the state row for this user (or the install-default NULL row).
  42. Creates the row lazily on first access so the migration doesn't need to
  43. seed anything.
  44. """
  45. if user_id is None:
  46. stmt = select(SponsorToastState).where(SponsorToastState.user_id.is_(None))
  47. else:
  48. stmt = select(SponsorToastState).where(SponsorToastState.user_id == user_id)
  49. result = await db.execute(stmt)
  50. state = result.scalar_one_or_none()
  51. if state is None:
  52. state = SponsorToastState(user_id=user_id, milestones_seen="[]")
  53. db.add(state)
  54. await db.flush()
  55. return state
  56. def _within_cooldown(state: SponsorToastState) -> bool:
  57. if state.last_shown_at is None:
  58. return False
  59. cutoff = datetime.now(timezone.utc) - timedelta(days=COOLDOWN_DAYS)
  60. last = state.last_shown_at
  61. if last.tzinfo is None:
  62. last = last.replace(tzinfo=timezone.utc)
  63. return last >= cutoff
  64. def _seen_milestones(state: SponsorToastState) -> set[str]:
  65. try:
  66. raw = json.loads(state.milestones_seen or "[]")
  67. return set(raw) if isinstance(raw, list) else set()
  68. except (json.JSONDecodeError, TypeError):
  69. logger.warning(
  70. "sponsor_toast_state.milestones_seen for user=%s was not valid JSON; resetting",
  71. state.user_id,
  72. )
  73. return set()
  74. # ---------------------------------------------------------------------------
  75. # Per-family checks
  76. # ---------------------------------------------------------------------------
  77. def _user_filter(column, user_id: int | None):
  78. return column.is_(None) if user_id is None else column == user_id
  79. async def _check_anniversary(
  80. db: AsyncSession, user_id: int | None, seen: set[str], _state: SponsorToastState
  81. ) -> Trigger | None:
  82. milestone = f"anniversary-{ANNIVERSARY_YEARS}"
  83. if milestone in seen:
  84. return None
  85. if user_id is None:
  86. # Install-anchor = earliest users.created_at (the first admin row).
  87. result = await db.execute(select(func.min(User.created_at)))
  88. anchor = result.scalar()
  89. else:
  90. result = await db.execute(select(User.created_at).where(User.id == user_id))
  91. anchor = result.scalar()
  92. if anchor is None:
  93. return None
  94. if anchor.tzinfo is None:
  95. anchor = anchor.replace(tzinfo=timezone.utc)
  96. if datetime.now(timezone.utc) - anchor < timedelta(days=365 * ANNIVERSARY_YEARS):
  97. return None
  98. return Trigger(milestone=milestone, family="anniversary")
  99. async def _check_prints(
  100. db: AsyncSession, user_id: int | None, seen: set[str], _state: SponsorToastState
  101. ) -> Trigger | None:
  102. stmt = (
  103. select(func.count())
  104. .select_from(PrintLogEntry)
  105. .where(
  106. PrintLogEntry.status == "completed",
  107. _user_filter(PrintLogEntry.created_by_id, user_id),
  108. )
  109. )
  110. completed = (await db.execute(stmt)).scalar() or 0
  111. # Pick the LARGEST milestone the user has crossed but not yet seen.
  112. for threshold in sorted(PRINT_MILESTONES, reverse=True):
  113. key = f"prints-{threshold}"
  114. if completed >= threshold and key not in seen:
  115. return Trigger(
  116. milestone=key,
  117. family="prints",
  118. threshold=threshold,
  119. payload={"count": completed},
  120. )
  121. return None
  122. async def _check_archives(
  123. db: AsyncSession, user_id: int | None, seen: set[str], _state: SponsorToastState
  124. ) -> Trigger | None:
  125. stmt = select(func.count()).select_from(PrintArchive).where(_user_filter(PrintArchive.created_by_id, user_id))
  126. archived = (await db.execute(stmt)).scalar() or 0
  127. for threshold in sorted(ARCHIVE_MILESTONES, reverse=True):
  128. key = f"archives-{threshold}"
  129. if archived >= threshold and key not in seen:
  130. return Trigger(
  131. milestone=key,
  132. family="archives",
  133. threshold=threshold,
  134. payload={"count": archived},
  135. )
  136. return None
  137. async def _check_cost(
  138. db: AsyncSession, user_id: int | None, seen: set[str], _state: SponsorToastState
  139. ) -> Trigger | None:
  140. stmt = (
  141. select(func.coalesce(func.sum(PrintLogEntry.cost), 0) + func.coalesce(func.sum(PrintLogEntry.energy_cost), 0))
  142. .select_from(PrintLogEntry)
  143. .where(_user_filter(PrintLogEntry.created_by_id, user_id))
  144. )
  145. total = float((await db.execute(stmt)).scalar() or 0)
  146. for threshold in sorted(COST_MILESTONES, reverse=True):
  147. key = f"cost-{threshold}"
  148. if total >= threshold and key not in seen:
  149. return Trigger(
  150. milestone=key,
  151. family="cost",
  152. threshold=threshold,
  153. payload={"total": round(total, 2)},
  154. )
  155. return None
  156. async def _check_version_update(
  157. _db: AsyncSession, _user_id: int | None, _seen: set[str], state: SponsorToastState
  158. ) -> Trigger | None:
  159. # version-update is NOT in milestones_seen — it has its own state column
  160. # so it can re-fire on each major bump.
  161. if not APP_VERSION:
  162. return None
  163. last = state.last_seen_version
  164. if last is None:
  165. # First-ever read; treat as already-acknowledged so we don't toast
  166. # immediately on a brand-new install. Persist current version silently.
  167. state.last_seen_version = APP_VERSION
  168. return None
  169. if last == APP_VERSION:
  170. return None
  171. return Trigger(
  172. milestone="version-update",
  173. family="version-update",
  174. payload={"from": last, "to": APP_VERSION},
  175. )
  176. # Priority order: most emotional / earned first; version-update is the soft fallback.
  177. _CHECKS = (
  178. _check_anniversary,
  179. _check_prints,
  180. _check_archives,
  181. _check_cost,
  182. _check_version_update,
  183. )
  184. # ---------------------------------------------------------------------------
  185. # Public API
  186. # ---------------------------------------------------------------------------
  187. async def evaluate(db: AsyncSession, user_id: int | None) -> Trigger | None:
  188. """Return the next eligible sponsor-toast trigger, or None."""
  189. state = await _get_or_create_state(db, user_id)
  190. if _within_cooldown(state):
  191. return None
  192. seen = _seen_milestones(state)
  193. for check in _CHECKS:
  194. trigger = await check(db, user_id, seen, state)
  195. if trigger is not None:
  196. return trigger
  197. # No triggers eligible — still commit any in-progress state changes
  198. # (e.g. version-update's first-touch persistence).
  199. await db.flush()
  200. return None
  201. async def dismiss(db: AsyncSession, user_id: int | None, milestone: str) -> None:
  202. """Mark a milestone as shown (sets cooldown anchor + records seen)."""
  203. state = await _get_or_create_state(db, user_id)
  204. if milestone == "version-update":
  205. # Re-armable: just update last_seen_version, don't add to seen-list.
  206. state.last_seen_version = APP_VERSION
  207. else:
  208. seen = _seen_milestones(state)
  209. if milestone not in seen:
  210. seen.add(milestone)
  211. state.milestones_seen = json.dumps(sorted(seen))
  212. state.last_shown_at = datetime.now(timezone.utc)
  213. await db.flush()