announcements.py 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468469470471472473474475476
  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. published_at = _parse_time(raw.get("published_at"))
  205. expires_at = _parse_time(raw.get("expires_at"))
  206. # History is whatever has expired; the feed's ``archived`` flag says the same
  207. # thing for messages the registrar kept after their expiry. One flagged but
  208. # still in date by this install's clock is history anyway: expire it now.
  209. now = datetime.now(timezone.utc).replace(tzinfo=None)
  210. if raw.get("archived") is True and (expires_at is None or expires_at > now):
  211. expires_at = now
  212. return Entry(
  213. public_id=public_id,
  214. level=level,
  215. texts=texts,
  216. link_url=link,
  217. published_at=published_at,
  218. expires_at=expires_at,
  219. )
  220. # ---- store ---------------------------------------------------------------------------
  221. async def _get(db: AsyncSession, key: str) -> str | None:
  222. return (await db.execute(select(Settings.value).where(Settings.key == key))).scalar_one_or_none()
  223. async def _set(db: AsyncSession, key: str, value: str) -> None:
  224. from backend.app.core.db_dialect import upsert_setting
  225. await upsert_setting(db, Settings, key, value)
  226. async def is_enabled(db: AsyncSession) -> bool:
  227. return (await _get(db, ENABLED_KEY) or "true").lower() != "false"
  228. async def apply_payload(db: AsyncSession, payload: dict, facts: InstallFacts) -> int:
  229. """Replace the stored announcements with the ones in ``payload`` that apply here.
  230. Refuses (FeedRejected) a serial below the highest already accepted. The same
  231. serial again is the same feed fetched twice, and is fine. Returns how many are
  232. stored. The caller commits.
  233. """
  234. serial = payload["serial"]
  235. seen = int(await _get(db, SERIAL_KEY) or 0)
  236. if serial < seen:
  237. raise FeedRejected(f"serial {serial} is older than {seen}, already accepted")
  238. wanted: dict[str, Entry] = {}
  239. for raw in payload["announcements"]:
  240. entry = parse_entry(raw)
  241. if entry is not None and targets(raw, facts):
  242. wanted[entry.public_id] = entry
  243. existing = {a.public_id: a for a in (await db.execute(select(Announcement))).scalars().all()}
  244. gone = [a.id for pid, a in existing.items() if pid not in wanted]
  245. if gone:
  246. # Read markers first: SQLite enforces no foreign keys unless asked to, so
  247. # ON DELETE CASCADE alone would leave them behind.
  248. await db.execute(delete(AnnouncementRead).where(AnnouncementRead.announcement_id.in_(gone)))
  249. await db.execute(delete(Announcement).where(Announcement.id.in_(gone)))
  250. for pid, entry in wanted.items():
  251. row = existing.get(pid) or Announcement(public_id=pid)
  252. row.level = entry.level
  253. row.texts = json.dumps(entry.texts, ensure_ascii=False)
  254. row.link_url = entry.link_url
  255. row.published_at = entry.published_at
  256. row.expires_at = entry.expires_at
  257. if pid not in existing:
  258. db.add(row)
  259. await _set(db, SERIAL_KEY, str(max(serial, seen)))
  260. return len(wanted)
  261. async def _download() -> bytes:
  262. # No version and no install identity in the request: the generic agent says
  263. # what is asking, nothing about which install.
  264. headers = {"User-Agent": "Bambuddy-Announcements", "Cache-Control": "no-cache"}
  265. async with (
  266. httpx.AsyncClient(timeout=15, follow_redirects=False) as client,
  267. client.stream("GET", FEED_URL, headers=headers) as response,
  268. ):
  269. if response.status_code != 200:
  270. raise FeedRejected(f"HTTP {response.status_code}")
  271. chunks: list[bytes] = []
  272. size = 0
  273. async for chunk in response.aiter_bytes():
  274. size += len(chunk)
  275. if size > MAX_FEED_BYTES:
  276. raise FeedRejected("larger than the feed is allowed to be")
  277. chunks.append(chunk)
  278. return b"".join(chunks)
  279. async def refresh(db: AsyncSession) -> int | None:
  280. """Fetch, verify and store. Returns how many apply here, or None if nothing changed.
  281. Never raises for a network problem or a bad feed -- those are logged and the
  282. last good list stays. A switched-off install fetches nothing.
  283. """
  284. if not await is_enabled(db):
  285. return None
  286. seen_before = int(await _get(db, SERIAL_KEY) or 0)
  287. try:
  288. content = await _download()
  289. payload = verify_feed(content)
  290. count = await apply_payload(db, payload, await install_facts(db))
  291. except FeedRejected as exc:
  292. await db.rollback()
  293. logger.warning("Announcements feed ignored: %s", exc)
  294. return None
  295. except httpx.HTTPError as exc:
  296. await db.rollback()
  297. # Offline and air-gapped installs land here every time; not worth a warning.
  298. logger.info("Announcements feed not fetched: %s", type(exc).__name__)
  299. return None
  300. await _set(db, LAST_FETCH_KEY, datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"))
  301. await db.commit()
  302. logger.info("Announcements feed #%s: %d for this install", payload["serial"], count)
  303. if payload["serial"] > seen_before:
  304. await _tell_open_pages()
  305. return count
  306. async def _tell_open_pages() -> None:
  307. """Have every open Bambuddy page re-read the list, so a new message's dot and
  308. banner appear without a reload. The event carries nothing: each page asks
  309. GET /announcements, which answers by who is asking."""
  310. from backend.app.core.websocket import ws_manager
  311. try:
  312. await ws_manager.broadcast({"type": "announcements_changed"})
  313. except Exception: # A page that misses it catches up on its own poll.
  314. logger.debug("announcements_changed broadcast failed", exc_info=True)
  315. # ---- read state ----------------------------------------------------------------------
  316. async def list_for(db: AsyncSession, user_id: int | None) -> list[dict]:
  317. """Every stored announcement, newest first, with this user's read state.
  318. ``archived`` marks history: messages past their expiry, which the panel lists
  319. under "Earlier" and which never count as unread or raise a banner. The feed
  320. decides how much history there is; a message it drops is deleted here.
  321. """
  322. now = datetime.now(timezone.utc).replace(tzinfo=None)
  323. rows = (
  324. (await db.execute(select(Announcement).order_by(Announcement.published_at.desc(), Announcement.id.desc())))
  325. .scalars()
  326. .all()
  327. )
  328. reader = AnnouncementRead.user_id.is_(None) if user_id is None else AnnouncementRead.user_id == user_id
  329. read_ids = set((await db.execute(select(AnnouncementRead.announcement_id).where(reader))).scalars().all())
  330. result = []
  331. for a in rows:
  332. try:
  333. texts = json.loads(a.texts)
  334. except (TypeError, ValueError):
  335. continue
  336. result.append(
  337. {
  338. "id": a.public_id,
  339. "level": a.level,
  340. "texts": texts,
  341. "link_url": a.link_url,
  342. "published_at": a.published_at.isoformat() + "Z" if a.published_at else None,
  343. "expires_at": a.expires_at.isoformat() + "Z" if a.expires_at else None,
  344. "archived": a.expires_at is not None and a.expires_at <= now,
  345. "read": a.id in read_ids,
  346. }
  347. )
  348. return result
  349. async def mark_read(db: AsyncSession, public_id: str, user_id: int | None) -> bool:
  350. """Record that this user read it. False if there is no such announcement."""
  351. announcement = (
  352. await db.execute(select(Announcement).where(Announcement.public_id == public_id))
  353. ).scalar_one_or_none()
  354. if announcement is None:
  355. return False
  356. reader = AnnouncementRead.user_id.is_(None) if user_id is None else AnnouncementRead.user_id == user_id
  357. # Checked rather than left to the unique constraint: NULL user_ids never
  358. # collide in one, so the auth-off row would be duplicated on every click.
  359. already = (
  360. await db.execute(select(AnnouncementRead.id).where(AnnouncementRead.announcement_id == announcement.id, reader))
  361. ).first()
  362. if already is None:
  363. db.add(AnnouncementRead(announcement_id=announcement.id, user_id=user_id))
  364. return True
  365. # ---- background loop -----------------------------------------------------------------
  366. _task: asyncio.Task | None = None
  367. async def _loop() -> None:
  368. from backend.app.core.database import async_session
  369. await asyncio.sleep(FIRST_FETCH_DELAY_SECONDS)
  370. while True:
  371. try:
  372. async with async_session() as db:
  373. await refresh(db)
  374. except asyncio.CancelledError:
  375. raise
  376. except Exception:
  377. logger.exception("Announcements refresh failed")
  378. await asyncio.sleep(FETCH_INTERVAL_SECONDS + random.uniform(0, FETCH_JITTER_SECONDS))
  379. def start() -> None:
  380. global _task
  381. if _task is None:
  382. _task = asyncio.create_task(_loop())
  383. def stop() -> None:
  384. global _task
  385. if _task is not None:
  386. _task.cancel()
  387. _task = None