announcements.py 17 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457
  1. """Announcements from the Bambuddy maintainers: fetch, verify, keep what applies.
  2. Bambuddy fetches one file, ``feed.json``, from the public
  3. ``maziggy/bambuddy-notifications`` repo on raw.githubusercontent.com -- the host
  4. the update check already talks to. No Bambuddy server is contacted and nothing
  5. about this install is sent: whether a message applies here (version range,
  6. channel, install type) is decided below, locally.
  7. The file is written by the maintainers' registrar and signed with Ed25519:
  8. {"format": 1, "key_id": "...", "signature": "<base64>", "payload": {...}}
  9. The signature covers ``canonical(payload)``. A file that does not verify against
  10. a key in ``TRUSTED_KEYS`` is ignored, so neither a copy of the repo nor anyone in
  11. the middle can make Bambuddy show a message. The payload's ``serial`` only goes
  12. up; a feed older than one already accepted is refused, so an old signed file
  13. cannot be re-served to bring back a withdrawn message.
  14. The feed is the full current list. On every accepted fetch the stored set is
  15. replaced: a message withdrawn upstream disappears here with its read markers.
  16. Any failure keeps the last good list.
  17. """
  18. from __future__ import annotations
  19. import asyncio
  20. import base64
  21. import json
  22. import logging
  23. import random
  24. import re
  25. from dataclasses import dataclass
  26. from datetime import datetime, timezone
  27. from urllib.parse import urlsplit
  28. import httpx
  29. from cryptography.exceptions import InvalidSignature
  30. from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey
  31. from sqlalchemy import delete, select
  32. from sqlalchemy.ext.asyncio import AsyncSession
  33. from backend.app.core.config import APP_VERSION
  34. from backend.app.models.announcement import Announcement, AnnouncementRead
  35. from backend.app.models.settings import Settings
  36. logger = logging.getLogger(__name__)
  37. FEED_URL = "https://raw.githubusercontent.com/maziggy/bambuddy-notifications/main/feed.json"
  38. FEED_FORMAT = 1
  39. # key_id (first 16 hex of sha256 of the raw public key) -> base64 raw public key.
  40. # A list so the key can be rotated: ship the new key next to the old one first.
  41. TRUSTED_KEYS: dict[str, str] = {
  42. "d70b3bf207fdfd59": "uQqUAYrInmXQ1nIwOIY/95L7tD5o/HZMmZfhAlULJ/w=",
  43. }
  44. # Where a message may link to. The registrar refuses anything else too; this is
  45. # the side that counts.
  46. LINK_HOSTS = ("github.com", "bambuddy.cool")
  47. LEVELS = ("info", "important", "critical")
  48. MAX_FEED_BYTES = 512 * 1024
  49. MAX_TITLE = 120
  50. MAX_BODY = 2000
  51. MAX_LINK_LABEL = 40
  52. FETCH_INTERVAL_SECONDS = 6 * 3600
  53. FETCH_JITTER_SECONDS = 30 * 60
  54. # Let startup settle before the first fetch.
  55. FIRST_FETCH_DELAY_SECONDS = 60
  56. ENABLED_KEY = "announcements_enabled"
  57. ALL_USERS_KEY = "announcements_all_users"
  58. # Internal state, never part of the settings API.
  59. SERIAL_KEY = "announcements_feed_serial"
  60. LAST_FETCH_KEY = "announcements_last_fetch"
  61. class FeedRejected(Exception):
  62. """The fetched file is not a feed this install accepts. Nothing changes."""
  63. # ---- verification -------------------------------------------------------------------
  64. def canonical(payload: dict) -> bytes:
  65. return json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode()
  66. def verify_feed(content: bytes, trusted_keys: dict[str, str] | None = None) -> dict:
  67. """The payload of a correctly signed feed, or FeedRejected."""
  68. keys = TRUSTED_KEYS if trusted_keys is None else trusted_keys
  69. try:
  70. envelope = json.loads(content)
  71. except (ValueError, UnicodeDecodeError) as exc:
  72. raise FeedRejected("not JSON") from exc
  73. if not isinstance(envelope, dict) or envelope.get("format") != FEED_FORMAT:
  74. raise FeedRejected("unknown feed format")
  75. public = keys.get(str(envelope.get("key_id")))
  76. if public is None:
  77. raise FeedRejected(f"signed with an unknown key ({envelope.get('key_id')!r})")
  78. payload = envelope.get("payload")
  79. if not isinstance(payload, dict):
  80. raise FeedRejected("no payload")
  81. try:
  82. signature = base64.b64decode(str(envelope.get("signature")), validate=True)
  83. Ed25519PublicKey.from_public_bytes(base64.b64decode(public)).verify(signature, canonical(payload))
  84. except (ValueError, InvalidSignature) as exc:
  85. raise FeedRejected("signature does not verify") from exc
  86. if payload.get("format") != FEED_FORMAT:
  87. raise FeedRejected("unknown payload format")
  88. serial = payload.get("serial")
  89. if not isinstance(serial, int) or isinstance(serial, bool) or serial < 1:
  90. raise FeedRejected("no serial")
  91. if not isinstance(payload.get("announcements"), list):
  92. raise FeedRejected("no announcement list")
  93. return payload
  94. # ---- what this install is ------------------------------------------------------------
  95. @dataclass(frozen=True)
  96. class InstallFacts:
  97. version: str
  98. channel: str # stable | beta
  99. install_type: str # docker | native | ha_addon | windows
  100. def _version_key(version: str) -> tuple:
  101. """Sortable form that puts a release above its own betas (1.2.6 > 1.2.6b3)."""
  102. from backend.app.api.routes.updates import parse_version
  103. parsed = parse_version(version)
  104. major, minor, patch, micro = (parsed + (0, 0, 0, 0))[:4]
  105. is_prerelease = parsed[4] if len(parsed) > 4 else 0
  106. prerelease_num = parsed[5] if len(parsed) > 5 else 0
  107. return (major, minor, patch, micro, 1 - is_prerelease, prerelease_num)
  108. def _install_type() -> str:
  109. from backend.app.api.routes import updates
  110. if updates._is_windows_installer_install():
  111. return "windows"
  112. if updates._is_ha_addon():
  113. return "ha_addon"
  114. if updates._is_docker_environment():
  115. return "docker"
  116. return "native"
  117. async def install_facts(db: AsyncSession) -> InstallFacts:
  118. beta_setting = (
  119. await db.execute(select(Settings.value).where(Settings.key == "include_beta_updates"))
  120. ).scalar_one_or_none()
  121. # Beta testers are the installs that asked for betas, and the ones running one.
  122. prerelease = bool(re.search(r"[a-zA-Z]", APP_VERSION.lstrip("v")))
  123. beta = prerelease or (beta_setting or "").lower() == "true"
  124. return InstallFacts(version=APP_VERSION, channel="beta" if beta else "stable", install_type=_install_type())
  125. def targets(entry: dict, facts: InstallFacts) -> bool:
  126. target = entry.get("target") or {}
  127. if not isinstance(target, dict):
  128. return False
  129. try:
  130. here = _version_key(facts.version)
  131. low, high = target.get("min_version"), target.get("max_version")
  132. if low and here < _version_key(str(low)):
  133. return False
  134. if high and here > _version_key(str(high)):
  135. return False
  136. except (TypeError, ValueError):
  137. return False
  138. channels = target.get("channels") or []
  139. if channels and facts.channel not in channels:
  140. return False
  141. install_types = target.get("install_types") or []
  142. return not (install_types and facts.install_type not in install_types)
  143. # ---- one entry, defensively ---------------------------------------------------------
  144. def link_allowed(url: str) -> bool:
  145. try:
  146. parts = urlsplit(url)
  147. port = parts.port
  148. except ValueError:
  149. return False
  150. host = (parts.hostname or "").lower()
  151. if parts.scheme != "https" or not host or parts.username or parts.password:
  152. return False
  153. if port not in (None, 443):
  154. return False
  155. return any(host == allowed or host.endswith("." + allowed) for allowed in LINK_HOSTS)
  156. def _parse_time(value: object) -> datetime | None:
  157. if not isinstance(value, str) or not value:
  158. return None
  159. try:
  160. parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
  161. except ValueError:
  162. return None
  163. if parsed.tzinfo is not None:
  164. parsed = parsed.astimezone(timezone.utc).replace(tzinfo=None)
  165. return parsed
  166. def _clean_texts(texts: object) -> dict[str, dict[str, str]] | None:
  167. """Plain strings, cut to length. None when there is no usable English text."""
  168. if not isinstance(texts, dict):
  169. return None
  170. cleaned: dict[str, dict[str, str]] = {}
  171. for lang, text in texts.items():
  172. if not isinstance(lang, str) or len(lang) > 10 or not isinstance(text, dict):
  173. continue
  174. title, body = text.get("title"), text.get("body")
  175. if not isinstance(title, str) or not isinstance(body, str) or not title.strip() or not body.strip():
  176. continue
  177. entry = {"title": title.strip()[:MAX_TITLE], "body": body.strip()[:MAX_BODY]}
  178. label = text.get("link_label")
  179. if isinstance(label, str) and label.strip():
  180. entry["link_label"] = label.strip()[:MAX_LINK_LABEL]
  181. cleaned[lang] = entry
  182. return cleaned if "en" in cleaned else None
  183. @dataclass
  184. class Entry:
  185. public_id: str
  186. level: str
  187. texts: dict[str, dict[str, str]]
  188. link_url: str | None
  189. published_at: datetime | None
  190. expires_at: datetime | None
  191. def parse_entry(raw: object) -> Entry | None:
  192. """One feed entry, or None if it can't be shown safely. Never raises."""
  193. if not isinstance(raw, dict):
  194. return None
  195. public_id = raw.get("id")
  196. if not isinstance(public_id, str) or not re.fullmatch(r"[A-Za-z0-9_-]{1,64}", public_id):
  197. return None
  198. level = raw.get("level") if raw.get("level") in LEVELS else "info"
  199. texts = _clean_texts(raw.get("texts"))
  200. if texts is None:
  201. return None
  202. link = raw.get("link_url")
  203. link = link if isinstance(link, str) and link_allowed(link) else None
  204. return Entry(
  205. public_id=public_id,
  206. level=level,
  207. texts=texts,
  208. link_url=link,
  209. published_at=_parse_time(raw.get("published_at")),
  210. expires_at=_parse_time(raw.get("expires_at")),
  211. )
  212. # ---- store ---------------------------------------------------------------------------
  213. async def _get(db: AsyncSession, key: str) -> str | None:
  214. return (await db.execute(select(Settings.value).where(Settings.key == key))).scalar_one_or_none()
  215. async def _set(db: AsyncSession, key: str, value: str) -> None:
  216. from backend.app.core.db_dialect import upsert_setting
  217. await upsert_setting(db, Settings, key, value)
  218. async def is_enabled(db: AsyncSession) -> bool:
  219. return (await _get(db, ENABLED_KEY) or "true").lower() != "false"
  220. async def apply_payload(db: AsyncSession, payload: dict, facts: InstallFacts) -> int:
  221. """Replace the stored announcements with the ones in ``payload`` that apply here.
  222. Refuses (FeedRejected) a serial below the highest already accepted. The same
  223. serial again is the same feed fetched twice, and is fine. Returns how many are
  224. stored. The caller commits.
  225. """
  226. serial = payload["serial"]
  227. seen = int(await _get(db, SERIAL_KEY) or 0)
  228. if serial < seen:
  229. raise FeedRejected(f"serial {serial} is older than {seen}, already accepted")
  230. wanted: dict[str, Entry] = {}
  231. for raw in payload["announcements"]:
  232. entry = parse_entry(raw)
  233. if entry is not None and targets(raw, facts):
  234. wanted[entry.public_id] = entry
  235. existing = {a.public_id: a for a in (await db.execute(select(Announcement))).scalars().all()}
  236. gone = [a.id for pid, a in existing.items() if pid not in wanted]
  237. if gone:
  238. # Read markers first: SQLite enforces no foreign keys unless asked to, so
  239. # ON DELETE CASCADE alone would leave them behind.
  240. await db.execute(delete(AnnouncementRead).where(AnnouncementRead.announcement_id.in_(gone)))
  241. await db.execute(delete(Announcement).where(Announcement.id.in_(gone)))
  242. for pid, entry in wanted.items():
  243. row = existing.get(pid) or Announcement(public_id=pid)
  244. row.level = entry.level
  245. row.texts = json.dumps(entry.texts, ensure_ascii=False)
  246. row.link_url = entry.link_url
  247. row.published_at = entry.published_at
  248. row.expires_at = entry.expires_at
  249. if pid not in existing:
  250. db.add(row)
  251. await _set(db, SERIAL_KEY, str(max(serial, seen)))
  252. return len(wanted)
  253. async def _download() -> bytes:
  254. # No version and no install identity in the request: the generic agent says
  255. # what is asking, nothing about which install.
  256. headers = {"User-Agent": "Bambuddy-Announcements", "Cache-Control": "no-cache"}
  257. async with (
  258. httpx.AsyncClient(timeout=15, follow_redirects=False) as client,
  259. client.stream("GET", FEED_URL, headers=headers) as response,
  260. ):
  261. if response.status_code != 200:
  262. raise FeedRejected(f"HTTP {response.status_code}")
  263. chunks: list[bytes] = []
  264. size = 0
  265. async for chunk in response.aiter_bytes():
  266. size += len(chunk)
  267. if size > MAX_FEED_BYTES:
  268. raise FeedRejected("larger than the feed is allowed to be")
  269. chunks.append(chunk)
  270. return b"".join(chunks)
  271. async def refresh(db: AsyncSession) -> int | None:
  272. """Fetch, verify and store. Returns how many apply here, or None if nothing changed.
  273. Never raises for a network problem or a bad feed -- those are logged and the
  274. last good list stays. A switched-off install fetches nothing.
  275. """
  276. if not await is_enabled(db):
  277. return None
  278. try:
  279. content = await _download()
  280. payload = verify_feed(content)
  281. count = await apply_payload(db, payload, await install_facts(db))
  282. except FeedRejected as exc:
  283. await db.rollback()
  284. logger.warning("Announcements feed ignored: %s", exc)
  285. return None
  286. except httpx.HTTPError as exc:
  287. await db.rollback()
  288. # Offline and air-gapped installs land here every time; not worth a warning.
  289. logger.info("Announcements feed not fetched: %s", type(exc).__name__)
  290. return None
  291. await _set(db, LAST_FETCH_KEY, datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"))
  292. await db.commit()
  293. logger.info("Announcements feed #%s: %d for this install", payload["serial"], count)
  294. return count
  295. # ---- read state ----------------------------------------------------------------------
  296. def _visible_filter(now: datetime):
  297. return (Announcement.expires_at.is_(None)) | (Announcement.expires_at > now)
  298. async def list_for(db: AsyncSession, user_id: int | None) -> list[dict]:
  299. """Live announcements, newest first, each with whether this user has read it."""
  300. now = datetime.now(timezone.utc).replace(tzinfo=None)
  301. rows = (
  302. (
  303. await db.execute(
  304. select(Announcement)
  305. .where(_visible_filter(now))
  306. .order_by(Announcement.published_at.desc(), Announcement.id.desc())
  307. )
  308. )
  309. .scalars()
  310. .all()
  311. )
  312. reader = AnnouncementRead.user_id.is_(None) if user_id is None else AnnouncementRead.user_id == user_id
  313. read_ids = set((await db.execute(select(AnnouncementRead.announcement_id).where(reader))).scalars().all())
  314. result = []
  315. for a in rows:
  316. try:
  317. texts = json.loads(a.texts)
  318. except (TypeError, ValueError):
  319. continue
  320. result.append(
  321. {
  322. "id": a.public_id,
  323. "level": a.level,
  324. "texts": texts,
  325. "link_url": a.link_url,
  326. "published_at": a.published_at.isoformat() + "Z" if a.published_at else None,
  327. "expires_at": a.expires_at.isoformat() + "Z" if a.expires_at else None,
  328. "read": a.id in read_ids,
  329. }
  330. )
  331. return result
  332. async def mark_read(db: AsyncSession, public_id: str, user_id: int | None) -> bool:
  333. """Record that this user read it. False if there is no such announcement."""
  334. announcement = (
  335. await db.execute(select(Announcement).where(Announcement.public_id == public_id))
  336. ).scalar_one_or_none()
  337. if announcement is None:
  338. return False
  339. reader = AnnouncementRead.user_id.is_(None) if user_id is None else AnnouncementRead.user_id == user_id
  340. # Checked rather than left to the unique constraint: NULL user_ids never
  341. # collide in one, so the auth-off row would be duplicated on every click.
  342. already = (
  343. await db.execute(select(AnnouncementRead.id).where(AnnouncementRead.announcement_id == announcement.id, reader))
  344. ).first()
  345. if already is None:
  346. db.add(AnnouncementRead(announcement_id=announcement.id, user_id=user_id))
  347. return True
  348. # ---- background loop -----------------------------------------------------------------
  349. _task: asyncio.Task | None = None
  350. async def _loop() -> None:
  351. from backend.app.core.database import async_session
  352. await asyncio.sleep(FIRST_FETCH_DELAY_SECONDS)
  353. while True:
  354. try:
  355. async with async_session() as db:
  356. await refresh(db)
  357. except asyncio.CancelledError:
  358. raise
  359. except Exception:
  360. logger.exception("Announcements refresh failed")
  361. await asyncio.sleep(FETCH_INTERVAL_SECONDS + random.uniform(0, FETCH_JITTER_SECONDS))
  362. def start() -> None:
  363. global _task
  364. if _task is None:
  365. _task = asyncio.create_task(_loop())
  366. def stop() -> None:
  367. global _task
  368. if _task is not None:
  369. _task.cancel()
  370. _task = None