test_announcements_service.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307
  1. """Announcements: verifying the signed feed, and keeping only what applies here.
  2. The feed is signed by the maintainers' registrar. These tests sign their own
  3. feeds with a throwaway key passed in as the trusted one, and check the rules an
  4. install applies: a bad signature, an unknown key or an older serial changes
  5. nothing; a message that does not target this install is not stored; a withdrawn
  6. one disappears with its read markers.
  7. """
  8. import base64
  9. import hashlib
  10. import json
  11. from datetime import datetime, timedelta, timezone
  12. from unittest.mock import AsyncMock, patch
  13. import pytest
  14. from cryptography.hazmat.primitives import serialization
  15. from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
  16. from sqlalchemy import select
  17. from backend.app.models.announcement import Announcement, AnnouncementRead
  18. from backend.app.models.settings import Settings
  19. from backend.app.services import announcements as svc
  20. from backend.app.services.announcements import FeedRejected, InstallFacts
  21. KEY = Ed25519PrivateKey.generate()
  22. _PUB = KEY.public_key().public_bytes(serialization.Encoding.Raw, serialization.PublicFormat.Raw)
  23. KEY_ID = hashlib.sha256(_PUB).hexdigest()[:16]
  24. TRUSTED = {KEY_ID: base64.b64encode(_PUB).decode()}
  25. STABLE_DOCKER = InstallFacts(version="1.2.6", channel="stable", install_type="docker")
  26. def entry(public_id="a1", level="info", target=None, **extra):
  27. raw = {
  28. "id": public_id,
  29. "level": level,
  30. "published_at": "2026-10-01T12:00:00Z",
  31. "expires_at": None,
  32. "texts": {"en": {"title": f"Title {public_id}", "body": "Body"}},
  33. "link_url": None,
  34. "target": target or {"min_version": None, "max_version": None, "channels": [], "install_types": []},
  35. }
  36. raw.update(extra)
  37. return raw
  38. def payload(*entries, serial=1):
  39. return {"format": 1, "serial": serial, "published_at": "2026-10-01T12:00:00Z", "announcements": list(entries)}
  40. def signed(p: dict, key: Ed25519PrivateKey = KEY, key_id: str = KEY_ID) -> bytes:
  41. """The file exactly as the registrar writes it: pretty, payload as an object."""
  42. sig = key.sign(svc.canonical(p))
  43. envelope = {"format": 1, "key_id": key_id, "signature": base64.b64encode(sig).decode(), "payload": p}
  44. return json.dumps(envelope, indent=2, ensure_ascii=False).encode()
  45. class TestVerify:
  46. def test_a_correctly_signed_feed_verifies(self):
  47. p = payload(entry(texts={"en": {"title": "Grüße — 你好", "body": "x"}}))
  48. assert svc.verify_feed(signed(p), TRUSTED) == p
  49. def test_tampered_text_is_refused(self):
  50. content = signed(payload(entry()))
  51. with pytest.raises(FeedRejected, match="signature"):
  52. svc.verify_feed(content.replace(b"Title a1", b"Title b1"), TRUSTED)
  53. def test_another_key_is_refused(self):
  54. stranger = Ed25519PrivateKey.generate()
  55. with pytest.raises(FeedRejected, match="signature"):
  56. svc.verify_feed(signed(payload(entry()), key=stranger), TRUSTED)
  57. def test_unknown_key_id_is_refused(self):
  58. with pytest.raises(FeedRejected, match="unknown key"):
  59. svc.verify_feed(signed(payload(entry()), key_id="0000000000000000"), TRUSTED)
  60. @pytest.mark.parametrize("content", [b"", b"not json", b"[]", b'{"format": 2}'])
  61. def test_garbage_is_refused(self, content):
  62. with pytest.raises(FeedRejected):
  63. svc.verify_feed(content, TRUSTED)
  64. @pytest.mark.parametrize("serial", [None, 0, "3", True])
  65. def test_a_payload_without_a_usable_serial_is_refused(self, serial):
  66. p = payload(entry())
  67. p["serial"] = serial
  68. with pytest.raises(FeedRejected, match="serial"):
  69. svc.verify_feed(signed(p), TRUSTED)
  70. def test_the_built_in_key_is_well_formed(self):
  71. for key_id, public in svc.TRUSTED_KEYS.items():
  72. raw = base64.b64decode(public)
  73. assert len(raw) == 32
  74. assert hashlib.sha256(raw).hexdigest()[:16] == key_id
  75. class TestEntries:
  76. @pytest.mark.parametrize("public_id", [None, "", "a b", "x" * 65, 7, "../x"])
  77. def test_unusable_ids_are_dropped(self, public_id):
  78. assert svc.parse_entry(entry(public_id=public_id)) is None
  79. def test_no_english_text_is_dropped(self):
  80. assert svc.parse_entry(entry(texts={"de": {"title": "Hallo", "body": "x"}})) is None
  81. def test_texts_are_cut_to_length_and_bad_languages_skipped(self):
  82. e = svc.parse_entry(
  83. entry(
  84. texts={
  85. "en": {"title": "t" * 500, "body": "b" * 5000, "link_label": "l" * 100},
  86. "de": {"title": "", "body": "x"},
  87. "fr": "not a dict",
  88. }
  89. )
  90. )
  91. assert set(e.texts) == {"en"}
  92. assert len(e.texts["en"]["title"]) == svc.MAX_TITLE
  93. assert len(e.texts["en"]["body"]) == svc.MAX_BODY
  94. assert len(e.texts["en"]["link_label"]) == svc.MAX_LINK_LABEL
  95. def test_unknown_level_is_info(self):
  96. assert svc.parse_entry(entry(level="apocalyptic")).level == "info"
  97. @pytest.mark.parametrize(
  98. "url",
  99. [
  100. "http://bambuddy.cool/",
  101. "https://evil.example/",
  102. "https://bambuddy.cool.evil.example/",
  103. "https://user@github.com/",
  104. "https://github.com:8443/",
  105. "javascript:alert(1)",
  106. ],
  107. )
  108. def test_links_off_the_allowlist_are_dropped_not_the_message(self, url):
  109. e = svc.parse_entry(entry(link_url=url))
  110. assert e is not None and e.link_url is None
  111. def test_allowed_link_is_kept(self):
  112. assert svc.parse_entry(entry(link_url="https://wiki.bambuddy.cool/x/")).link_url == (
  113. "https://wiki.bambuddy.cool/x/"
  114. )
  115. def test_times_are_naive_utc(self):
  116. e = svc.parse_entry(entry(expires_at="2026-11-01T02:00:00+02:00"))
  117. assert e.expires_at == datetime(2026, 11, 1, 0, 0)
  118. class TestTargeting:
  119. @pytest.mark.parametrize(
  120. "version, low, high, shown",
  121. [
  122. ("1.2.6", None, None, True),
  123. ("1.2.6", "1.2.6", None, True),
  124. ("1.2.6", "1.2.7", None, False),
  125. ("1.2.6", None, "1.2.5", False),
  126. ("1.2.6b1", "1.2.6", None, False), # a beta is older than its release
  127. ("1.2.6b1", None, "1.2.6", True),
  128. ("1.2.6b3", "1.2.6b2", "1.2.6b3", True),
  129. ],
  130. )
  131. def test_version_range(self, version, low, high, shown):
  132. facts = InstallFacts(version=version, channel="stable", install_type="docker")
  133. target = {"min_version": low, "max_version": high}
  134. assert svc.targets({"target": target}, facts) is shown
  135. def test_channel(self):
  136. beta_only = {"target": {"channels": ["beta"]}}
  137. assert not svc.targets(beta_only, STABLE_DOCKER)
  138. assert svc.targets(beta_only, InstallFacts("1.2.6", "beta", "docker"))
  139. def test_install_type(self):
  140. windows_only = {"target": {"install_types": ["windows"]}}
  141. assert not svc.targets(windows_only, STABLE_DOCKER)
  142. assert svc.targets(windows_only, InstallFacts("1.2.6", "stable", "windows"))
  143. def test_a_malformed_target_shows_nothing(self):
  144. assert not svc.targets({"target": "everyone"}, STABLE_DOCKER)
  145. async def _stored(db) -> list[str]:
  146. return sorted((await db.execute(select(Announcement.public_id))).scalars().all())
  147. async def _serial(db) -> str | None:
  148. return (await db.execute(select(Settings.value).where(Settings.key == svc.SERIAL_KEY))).scalar_one_or_none()
  149. class TestApply:
  150. @pytest.mark.asyncio
  151. async def test_stores_only_what_targets_this_install(self, db_session):
  152. p = payload(
  153. entry("everyone"),
  154. entry("beta", target={"channels": ["beta"]}),
  155. entry("windows", target={"install_types": ["windows"]}),
  156. entry("broken", texts={}),
  157. )
  158. assert await svc.apply_payload(db_session, p, STABLE_DOCKER) == 1
  159. await db_session.commit()
  160. assert await _stored(db_session) == ["everyone"]
  161. assert await _serial(db_session) == "1"
  162. @pytest.mark.asyncio
  163. async def test_withdrawn_upstream_goes_with_its_read_markers(self, db_session):
  164. await svc.apply_payload(db_session, payload(entry("keep"), entry("drop")), STABLE_DOCKER)
  165. await db_session.commit()
  166. assert await svc.mark_read(db_session, "drop", None)
  167. await db_session.commit()
  168. await svc.apply_payload(db_session, payload(entry("keep"), serial=2), STABLE_DOCKER)
  169. await db_session.commit()
  170. assert await _stored(db_session) == ["keep"]
  171. assert (await db_session.execute(select(AnnouncementRead))).scalars().all() == []
  172. @pytest.mark.asyncio
  173. async def test_an_edit_keeps_the_read_state(self, db_session):
  174. await svc.apply_payload(db_session, payload(entry("a1")), STABLE_DOCKER)
  175. await db_session.commit()
  176. await svc.mark_read(db_session, "a1", None)
  177. await db_session.commit()
  178. edited = entry("a1", texts={"en": {"title": "Fixed typo", "body": "Body"}})
  179. await svc.apply_payload(db_session, payload(edited, serial=2), STABLE_DOCKER)
  180. await db_session.commit()
  181. [item] = await svc.list_for(db_session, None)
  182. assert item["texts"]["en"]["title"] == "Fixed typo"
  183. assert item["read"] is True
  184. @pytest.mark.asyncio
  185. async def test_an_older_serial_is_refused_and_changes_nothing(self, db_session):
  186. await svc.apply_payload(db_session, payload(entry("new"), serial=5), STABLE_DOCKER)
  187. await db_session.commit()
  188. with pytest.raises(FeedRejected, match="older"):
  189. await svc.apply_payload(db_session, payload(entry("old"), serial=4), STABLE_DOCKER)
  190. await db_session.rollback()
  191. assert await _stored(db_session) == ["new"]
  192. assert await _serial(db_session) == "5"
  193. @pytest.mark.asyncio
  194. async def test_the_same_serial_again_is_fine(self, db_session):
  195. await svc.apply_payload(db_session, payload(entry("a1"), serial=3), STABLE_DOCKER)
  196. await svc.apply_payload(db_session, payload(entry("a1"), serial=3), STABLE_DOCKER)
  197. await db_session.commit()
  198. assert await _stored(db_session) == ["a1"]
  199. class TestReadState:
  200. @pytest.mark.asyncio
  201. async def test_expired_ones_are_not_listed(self, db_session):
  202. past = (datetime.now(timezone.utc) - timedelta(minutes=1)).strftime("%Y-%m-%dT%H:%M:%SZ")
  203. future = (datetime.now(timezone.utc) + timedelta(days=1)).strftime("%Y-%m-%dT%H:%M:%SZ")
  204. p = payload(entry("gone", expires_at=past), entry("live", expires_at=future), entry("forever"))
  205. await svc.apply_payload(db_session, p, STABLE_DOCKER)
  206. await db_session.commit()
  207. assert sorted(i["id"] for i in await svc.list_for(db_session, None)) == ["forever", "live"]
  208. @pytest.mark.asyncio
  209. async def test_reading_twice_with_auth_off_records_once(self, db_session):
  210. await svc.apply_payload(db_session, payload(entry("a1")), STABLE_DOCKER)
  211. await db_session.commit()
  212. assert await svc.mark_read(db_session, "a1", None)
  213. await db_session.commit()
  214. assert await svc.mark_read(db_session, "a1", None)
  215. await db_session.commit()
  216. assert len((await db_session.execute(select(AnnouncementRead))).scalars().all()) == 1
  217. @pytest.mark.asyncio
  218. async def test_unknown_id_is_not_marked(self, db_session):
  219. assert not await svc.mark_read(db_session, "nope", None)
  220. class TestRefresh:
  221. @pytest.mark.asyncio
  222. async def test_switched_off_fetches_nothing(self, db_session):
  223. db_session.add(Settings(key=svc.ENABLED_KEY, value="false"))
  224. await db_session.commit()
  225. with patch.object(svc, "_download", AsyncMock()) as download:
  226. assert await svc.refresh(db_session) is None
  227. download.assert_not_called()
  228. @pytest.mark.asyncio
  229. async def test_a_good_feed_is_stored(self, db_session):
  230. with (
  231. patch.object(svc, "_download", AsyncMock(return_value=signed(payload(entry("a1"))))),
  232. patch.object(svc, "TRUSTED_KEYS", TRUSTED),
  233. patch.object(svc, "install_facts", AsyncMock(return_value=STABLE_DOCKER)),
  234. ):
  235. assert await svc.refresh(db_session) == 1
  236. assert await _stored(db_session) == ["a1"]
  237. @pytest.mark.asyncio
  238. async def test_a_bad_feed_keeps_the_last_good_list(self, db_session):
  239. with (
  240. patch.object(svc, "TRUSTED_KEYS", TRUSTED),
  241. patch.object(svc, "install_facts", AsyncMock(return_value=STABLE_DOCKER)),
  242. ):
  243. with patch.object(svc, "_download", AsyncMock(return_value=signed(payload(entry("a1"))))):
  244. await svc.refresh(db_session)
  245. forged = signed(payload(serial=9), key=Ed25519PrivateKey.generate())
  246. with patch.object(svc, "_download", AsyncMock(return_value=forged)):
  247. assert await svc.refresh(db_session) is None
  248. assert await _stored(db_session) == ["a1"]
  249. @pytest.mark.asyncio
  250. async def test_offline_is_quiet_and_harmless(self, db_session):
  251. import httpx
  252. with patch.object(svc, "_download", AsyncMock(side_effect=httpx.ConnectError("offline"))):
  253. assert await svc.refresh(db_session) is None