test_sponsor_prompt_service.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327
  1. """Unit tests for the sponsor-prompt trigger evaluator."""
  2. from __future__ import annotations
  3. import json
  4. from datetime import datetime, timedelta, timezone
  5. from unittest.mock import patch
  6. import pytest
  7. from sqlalchemy.ext.asyncio import AsyncSession
  8. from backend.app.models.archive import PrintArchive
  9. from backend.app.models.print_log import PrintLogEntry
  10. from backend.app.models.sponsor_toast_state import SponsorToastState
  11. from backend.app.models.user import User
  12. from backend.app.services import sponsor_prompt as service
  13. # ---------------------------------------------------------------------------
  14. # Helpers
  15. # ---------------------------------------------------------------------------
  16. async def _make_user(db: AsyncSession, *, username: str = "alice", created_days_ago: int = 0) -> User:
  17. user = User(username=username, role="admin")
  18. db.add(user)
  19. await db.flush()
  20. if created_days_ago:
  21. user.created_at = datetime.now(timezone.utc) - timedelta(days=created_days_ago)
  22. await db.flush()
  23. return user
  24. async def _add_completed_prints(db: AsyncSession, *, user_id: int | None, count: int, cost_each: float = 0.0) -> None:
  25. for _ in range(count):
  26. db.add(
  27. PrintLogEntry(
  28. status="completed",
  29. created_by_id=user_id,
  30. cost=cost_each if cost_each else None,
  31. )
  32. )
  33. await db.flush()
  34. async def _add_archives(db: AsyncSession, *, user_id: int | None, count: int) -> None:
  35. for i in range(count):
  36. db.add(
  37. PrintArchive(
  38. filename=f"archive-{i}.zip",
  39. file_path=f"/tmp/archive-{i}.zip",
  40. file_size=1024,
  41. created_by_id=user_id,
  42. )
  43. )
  44. await db.flush()
  45. # ---------------------------------------------------------------------------
  46. # Empty / no-eligibility cases
  47. # ---------------------------------------------------------------------------
  48. class TestEmptyState:
  49. @pytest.mark.asyncio
  50. async def test_evaluate_returns_none_for_fresh_user(self, db_session: AsyncSession):
  51. user = await _make_user(db_session)
  52. trigger = await service.evaluate(db_session, user.id)
  53. assert trigger is None
  54. @pytest.mark.asyncio
  55. async def test_state_row_is_created_lazily(self, db_session: AsyncSession):
  56. user = await _make_user(db_session)
  57. await service.evaluate(db_session, user.id)
  58. from sqlalchemy import select
  59. row = (
  60. await db_session.execute(select(SponsorToastState).where(SponsorToastState.user_id == user.id))
  61. ).scalar_one_or_none()
  62. assert row is not None
  63. assert row.milestones_seen == "[]"
  64. # ---------------------------------------------------------------------------
  65. # Cooldown
  66. # ---------------------------------------------------------------------------
  67. class TestCooldown:
  68. @pytest.mark.asyncio
  69. async def test_no_toast_within_14d_window(self, db_session: AsyncSession):
  70. user = await _make_user(db_session)
  71. await _add_completed_prints(db_session, user_id=user.id, count=200)
  72. # Pre-populate state with a recent last_shown_at
  73. state = SponsorToastState(
  74. user_id=user.id,
  75. last_shown_at=datetime.now(timezone.utc) - timedelta(days=3),
  76. )
  77. db_session.add(state)
  78. await db_session.flush()
  79. trigger = await service.evaluate(db_session, user.id)
  80. assert trigger is None
  81. @pytest.mark.asyncio
  82. async def test_toast_eligible_after_14d_window(self, db_session: AsyncSession):
  83. user = await _make_user(db_session)
  84. await _add_completed_prints(db_session, user_id=user.id, count=200)
  85. state = SponsorToastState(
  86. user_id=user.id,
  87. last_shown_at=datetime.now(timezone.utc) - timedelta(days=15),
  88. )
  89. db_session.add(state)
  90. await db_session.flush()
  91. trigger = await service.evaluate(db_session, user.id)
  92. assert trigger is not None
  93. assert trigger.family == "prints"
  94. # ---------------------------------------------------------------------------
  95. # Per-family triggers
  96. # ---------------------------------------------------------------------------
  97. class TestPrintMilestones:
  98. @pytest.mark.asyncio
  99. async def test_fires_at_100(self, db_session: AsyncSession):
  100. user = await _make_user(db_session)
  101. await _add_completed_prints(db_session, user_id=user.id, count=100)
  102. trigger = await service.evaluate(db_session, user.id)
  103. assert trigger is not None
  104. assert trigger.milestone == "prints-100"
  105. assert trigger.threshold == 100
  106. @pytest.mark.asyncio
  107. async def test_picks_highest_unseen_milestone(self, db_session: AsyncSession):
  108. user = await _make_user(db_session)
  109. await _add_completed_prints(db_session, user_id=user.id, count=600)
  110. trigger = await service.evaluate(db_session, user.id)
  111. # 500 is the highest crossed milestone (1000 not reached).
  112. assert trigger is not None
  113. assert trigger.milestone == "prints-500"
  114. @pytest.mark.asyncio
  115. async def test_skips_already_seen(self, db_session: AsyncSession):
  116. user = await _make_user(db_session)
  117. await _add_completed_prints(db_session, user_id=user.id, count=600)
  118. # Mark prints-500 as already seen — but NOT prints-100.
  119. # Service should fall through to the next-largest unseen, which is prints-100.
  120. state = SponsorToastState(
  121. user_id=user.id,
  122. milestones_seen=json.dumps(["prints-500"]),
  123. )
  124. db_session.add(state)
  125. await db_session.flush()
  126. trigger = await service.evaluate(db_session, user.id)
  127. assert trigger is not None
  128. assert trigger.milestone == "prints-100"
  129. @pytest.mark.asyncio
  130. async def test_failed_prints_dont_count(self, db_session: AsyncSession):
  131. user = await _make_user(db_session)
  132. await _add_completed_prints(db_session, user_id=user.id, count=50)
  133. for _ in range(60):
  134. db_session.add(PrintLogEntry(status="failed", created_by_id=user.id))
  135. await db_session.flush()
  136. trigger = await service.evaluate(db_session, user.id)
  137. # Only 50 completed → below 100 threshold → no print trigger.
  138. # Anniversary not reached either; no other counter populated.
  139. assert trigger is None
  140. class TestArchiveMilestones:
  141. @pytest.mark.asyncio
  142. async def test_fires_at_50(self, db_session: AsyncSession):
  143. user = await _make_user(db_session)
  144. await _add_archives(db_session, user_id=user.id, count=50)
  145. trigger = await service.evaluate(db_session, user.id)
  146. assert trigger is not None
  147. assert trigger.milestone == "archives-50"
  148. class TestCostMilestones:
  149. @pytest.mark.asyncio
  150. async def test_fires_when_cost_sum_crosses_100(self, db_session: AsyncSession):
  151. # Prints with cost = ~3.5 each, 30 prints → 105.
  152. user = await _make_user(db_session)
  153. await _add_completed_prints(db_session, user_id=user.id, count=30, cost_each=3.5)
  154. # 30 < 100 prints, so prints-100 not eligible. cost = 105 ≥ 100 → fires.
  155. trigger = await service.evaluate(db_session, user.id)
  156. assert trigger is not None
  157. assert trigger.family == "cost"
  158. assert trigger.milestone == "cost-100"
  159. class TestAnniversary:
  160. @pytest.mark.asyncio
  161. async def test_fires_after_1_year(self, db_session: AsyncSession):
  162. user = await _make_user(db_session, created_days_ago=370)
  163. trigger = await service.evaluate(db_session, user.id)
  164. assert trigger is not None
  165. assert trigger.milestone == "anniversary-1"
  166. assert trigger.family == "anniversary"
  167. @pytest.mark.asyncio
  168. async def test_does_not_fire_before_1_year(self, db_session: AsyncSession):
  169. user = await _make_user(db_session, created_days_ago=300)
  170. trigger = await service.evaluate(db_session, user.id)
  171. assert trigger is None
  172. class TestVersionUpdate:
  173. @pytest.mark.asyncio
  174. async def test_first_read_silently_anchors(self, db_session: AsyncSession):
  175. user = await _make_user(db_session)
  176. with patch.object(service, "APP_VERSION", "0.3.0"):
  177. trigger = await service.evaluate(db_session, user.id)
  178. assert trigger is None
  179. from sqlalchemy import select
  180. state = (
  181. await db_session.execute(select(SponsorToastState).where(SponsorToastState.user_id == user.id))
  182. ).scalar_one()
  183. assert state.last_seen_version == "0.3.0"
  184. @pytest.mark.asyncio
  185. async def test_fires_on_version_bump(self, db_session: AsyncSession):
  186. user = await _make_user(db_session)
  187. state = SponsorToastState(user_id=user.id, last_seen_version="0.2.0")
  188. db_session.add(state)
  189. await db_session.flush()
  190. with patch.object(service, "APP_VERSION", "0.3.0"):
  191. trigger = await service.evaluate(db_session, user.id)
  192. assert trigger is not None
  193. assert trigger.milestone == "version-update"
  194. assert trigger.payload == {"from": "0.2.0", "to": "0.3.0"}
  195. # ---------------------------------------------------------------------------
  196. # Priority order
  197. # ---------------------------------------------------------------------------
  198. class TestPriorityOrder:
  199. @pytest.mark.asyncio
  200. async def test_anniversary_beats_prints(self, db_session: AsyncSession):
  201. # User old enough for anniversary AND with 100+ prints.
  202. user = await _make_user(db_session, created_days_ago=400)
  203. await _add_completed_prints(db_session, user_id=user.id, count=200)
  204. trigger = await service.evaluate(db_session, user.id)
  205. assert trigger is not None
  206. assert trigger.family == "anniversary"
  207. @pytest.mark.asyncio
  208. async def test_prints_beats_archives(self, db_session: AsyncSession):
  209. user = await _make_user(db_session)
  210. await _add_completed_prints(db_session, user_id=user.id, count=200)
  211. await _add_archives(db_session, user_id=user.id, count=100)
  212. trigger = await service.evaluate(db_session, user.id)
  213. assert trigger is not None
  214. assert trigger.family == "prints"
  215. # ---------------------------------------------------------------------------
  216. # Dismiss
  217. # ---------------------------------------------------------------------------
  218. class TestDismiss:
  219. @pytest.mark.asyncio
  220. async def test_dismiss_adds_to_seen_and_anchors_cooldown(self, db_session: AsyncSession):
  221. user = await _make_user(db_session)
  222. await _add_completed_prints(db_session, user_id=user.id, count=100)
  223. await service.evaluate(db_session, user.id)
  224. await service.dismiss(db_session, user.id, "prints-100")
  225. from sqlalchemy import select
  226. state = (
  227. await db_session.execute(select(SponsorToastState).where(SponsorToastState.user_id == user.id))
  228. ).scalar_one()
  229. assert "prints-100" in json.loads(state.milestones_seen)
  230. assert state.last_shown_at is not None
  231. # Re-evaluation must now return None (cooldown).
  232. next_trigger = await service.evaluate(db_session, user.id)
  233. assert next_trigger is None
  234. @pytest.mark.asyncio
  235. async def test_version_update_dismiss_updates_version_not_seen_list(self, db_session: AsyncSession):
  236. user = await _make_user(db_session)
  237. state = SponsorToastState(user_id=user.id, last_seen_version="0.2.0")
  238. db_session.add(state)
  239. await db_session.flush()
  240. with patch.object(service, "APP_VERSION", "0.3.0"):
  241. await service.dismiss(db_session, user.id, "version-update")
  242. from sqlalchemy import select
  243. state = (
  244. await db_session.execute(select(SponsorToastState).where(SponsorToastState.user_id == user.id))
  245. ).scalar_one()
  246. assert state.last_seen_version == "0.3.0"
  247. assert json.loads(state.milestones_seen) == []
  248. # ---------------------------------------------------------------------------
  249. # Auth-disabled (user_id = None) — NULL-keyed install-default row
  250. # ---------------------------------------------------------------------------
  251. class TestAuthDisabledMode:
  252. @pytest.mark.asyncio
  253. async def test_uses_install_anchor_for_anniversary(self, db_session: AsyncSession):
  254. # In auth-disabled mode, anniversary anchor = MIN(users.created_at).
  255. # Seed a user from >1 year ago.
  256. await _make_user(db_session, username="root", created_days_ago=400)
  257. # Prints written without created_by_id.
  258. await _add_completed_prints(db_session, user_id=None, count=10)
  259. trigger = await service.evaluate(db_session, None)
  260. assert trigger is not None
  261. assert trigger.family == "anniversary"
  262. @pytest.mark.asyncio
  263. async def test_null_keyed_counters_isolated_from_per_user(self, db_session: AsyncSession):
  264. # A user-attributed prints set should NOT show up in the install-default count.
  265. user = await _make_user(db_session, username="alice")
  266. await _add_completed_prints(db_session, user_id=user.id, count=200)
  267. # NULL-keyed install has zero prints.
  268. trigger = await service.evaluate(db_session, None)
  269. # No anniversary either (user only just created).
  270. assert trigger is None