test_sponsor_prompt_service.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336
  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", # nosec B108
  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_fires_at_lowest_threshold(self, db_session: AsyncSession):
  131. user = await _make_user(db_session)
  132. await _add_completed_prints(db_session, user_id=user.id, count=10)
  133. trigger = await service.evaluate(db_session, user.id)
  134. assert trigger is not None
  135. assert trigger.milestone == "prints-10"
  136. assert trigger.threshold == 10
  137. @pytest.mark.asyncio
  138. async def test_failed_prints_dont_count(self, db_session: AsyncSession):
  139. user = await _make_user(db_session)
  140. await _add_completed_prints(db_session, user_id=user.id, count=5)
  141. for _ in range(60):
  142. db_session.add(PrintLogEntry(status="failed", created_by_id=user.id))
  143. await db_session.flush()
  144. trigger = await service.evaluate(db_session, user.id)
  145. # Only 5 completed → below 10 threshold → no print trigger.
  146. # Anniversary not reached either; no other counter populated.
  147. assert trigger is None
  148. class TestArchiveMilestones:
  149. @pytest.mark.asyncio
  150. async def test_fires_at_50(self, db_session: AsyncSession):
  151. user = await _make_user(db_session)
  152. await _add_archives(db_session, user_id=user.id, count=50)
  153. trigger = await service.evaluate(db_session, user.id)
  154. assert trigger is not None
  155. assert trigger.milestone == "archives-50"
  156. class TestCostMilestones:
  157. @pytest.mark.asyncio
  158. async def test_fires_when_cost_sum_crosses_100(self, db_session: AsyncSession):
  159. # 5 prints, cost ~21 each → 105 total. Below the 10-print threshold so
  160. # the prints family stays silent and cost gets a chance.
  161. user = await _make_user(db_session)
  162. await _add_completed_prints(db_session, user_id=user.id, count=5, cost_each=21.0)
  163. trigger = await service.evaluate(db_session, user.id)
  164. assert trigger is not None
  165. assert trigger.family == "cost"
  166. assert trigger.milestone == "cost-100"
  167. class TestAnniversary:
  168. @pytest.mark.asyncio
  169. async def test_fires_after_1_year(self, db_session: AsyncSession):
  170. user = await _make_user(db_session, created_days_ago=370)
  171. trigger = await service.evaluate(db_session, user.id)
  172. assert trigger is not None
  173. assert trigger.milestone == "anniversary-1"
  174. assert trigger.family == "anniversary"
  175. @pytest.mark.asyncio
  176. async def test_does_not_fire_before_1_year(self, db_session: AsyncSession):
  177. user = await _make_user(db_session, created_days_ago=300)
  178. trigger = await service.evaluate(db_session, user.id)
  179. assert trigger is None
  180. class TestVersionUpdate:
  181. @pytest.mark.asyncio
  182. async def test_first_read_silently_anchors(self, db_session: AsyncSession):
  183. user = await _make_user(db_session)
  184. with patch.object(service, "APP_VERSION", "0.3.0"):
  185. trigger = await service.evaluate(db_session, user.id)
  186. assert trigger is None
  187. from sqlalchemy import select
  188. state = (
  189. await db_session.execute(select(SponsorToastState).where(SponsorToastState.user_id == user.id))
  190. ).scalar_one()
  191. assert state.last_seen_version == "0.3.0"
  192. @pytest.mark.asyncio
  193. async def test_fires_on_version_bump(self, db_session: AsyncSession):
  194. user = await _make_user(db_session)
  195. state = SponsorToastState(user_id=user.id, last_seen_version="0.2.0")
  196. db_session.add(state)
  197. await db_session.flush()
  198. with patch.object(service, "APP_VERSION", "0.3.0"):
  199. trigger = await service.evaluate(db_session, user.id)
  200. assert trigger is not None
  201. assert trigger.milestone == "version-update"
  202. assert trigger.payload == {"from": "0.2.0", "to": "0.3.0"}
  203. # ---------------------------------------------------------------------------
  204. # Priority order
  205. # ---------------------------------------------------------------------------
  206. class TestPriorityOrder:
  207. @pytest.mark.asyncio
  208. async def test_anniversary_beats_prints(self, db_session: AsyncSession):
  209. # User old enough for anniversary AND with 100+ prints.
  210. user = await _make_user(db_session, created_days_ago=400)
  211. await _add_completed_prints(db_session, user_id=user.id, count=200)
  212. trigger = await service.evaluate(db_session, user.id)
  213. assert trigger is not None
  214. assert trigger.family == "anniversary"
  215. @pytest.mark.asyncio
  216. async def test_prints_beats_archives(self, db_session: AsyncSession):
  217. user = await _make_user(db_session)
  218. await _add_completed_prints(db_session, user_id=user.id, count=200)
  219. await _add_archives(db_session, user_id=user.id, count=100)
  220. trigger = await service.evaluate(db_session, user.id)
  221. assert trigger is not None
  222. assert trigger.family == "prints"
  223. # ---------------------------------------------------------------------------
  224. # Dismiss
  225. # ---------------------------------------------------------------------------
  226. class TestDismiss:
  227. @pytest.mark.asyncio
  228. async def test_dismiss_adds_to_seen_and_anchors_cooldown(self, db_session: AsyncSession):
  229. user = await _make_user(db_session)
  230. await _add_completed_prints(db_session, user_id=user.id, count=100)
  231. await service.evaluate(db_session, user.id)
  232. await service.dismiss(db_session, user.id, "prints-100")
  233. from sqlalchemy import select
  234. state = (
  235. await db_session.execute(select(SponsorToastState).where(SponsorToastState.user_id == user.id))
  236. ).scalar_one()
  237. assert "prints-100" in json.loads(state.milestones_seen)
  238. assert state.last_shown_at is not None
  239. # Re-evaluation must now return None (cooldown).
  240. next_trigger = await service.evaluate(db_session, user.id)
  241. assert next_trigger is None
  242. @pytest.mark.asyncio
  243. async def test_version_update_dismiss_updates_version_not_seen_list(self, db_session: AsyncSession):
  244. user = await _make_user(db_session)
  245. state = SponsorToastState(user_id=user.id, last_seen_version="0.2.0")
  246. db_session.add(state)
  247. await db_session.flush()
  248. with patch.object(service, "APP_VERSION", "0.3.0"):
  249. await service.dismiss(db_session, user.id, "version-update")
  250. from sqlalchemy import select
  251. state = (
  252. await db_session.execute(select(SponsorToastState).where(SponsorToastState.user_id == user.id))
  253. ).scalar_one()
  254. assert state.last_seen_version == "0.3.0"
  255. assert json.loads(state.milestones_seen) == []
  256. # ---------------------------------------------------------------------------
  257. # Auth-disabled (user_id = None) — NULL-keyed install-default row
  258. # ---------------------------------------------------------------------------
  259. class TestAuthDisabledMode:
  260. @pytest.mark.asyncio
  261. async def test_uses_install_anchor_for_anniversary(self, db_session: AsyncSession):
  262. # In auth-disabled mode, anniversary anchor = MIN(users.created_at).
  263. # Seed a user from >1 year ago.
  264. await _make_user(db_session, username="root", created_days_ago=400)
  265. # Prints written without created_by_id.
  266. await _add_completed_prints(db_session, user_id=None, count=10)
  267. trigger = await service.evaluate(db_session, None)
  268. assert trigger is not None
  269. assert trigger.family == "anniversary"
  270. @pytest.mark.asyncio
  271. async def test_null_keyed_counters_isolated_from_per_user(self, db_session: AsyncSession):
  272. # A user-attributed prints set should NOT show up in the install-default count.
  273. user = await _make_user(db_session, username="alice")
  274. await _add_completed_prints(db_session, user_id=user.id, count=200)
  275. # NULL-keyed install has zero prints.
  276. trigger = await service.evaluate(db_session, None)
  277. # No anniversary either (user only just created).
  278. assert trigger is None