| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457 |
- """Announcements from the Bambuddy maintainers: fetch, verify, keep what applies.
- Bambuddy fetches one file, ``feed.json``, from the public
- ``maziggy/bambuddy-notifications`` repo on raw.githubusercontent.com -- the host
- the update check already talks to. No Bambuddy server is contacted and nothing
- about this install is sent: whether a message applies here (version range,
- channel, install type) is decided below, locally.
- The file is written by the maintainers' registrar and signed with Ed25519:
- {"format": 1, "key_id": "...", "signature": "<base64>", "payload": {...}}
- The signature covers ``canonical(payload)``. A file that does not verify against
- a key in ``TRUSTED_KEYS`` is ignored, so neither a copy of the repo nor anyone in
- the middle can make Bambuddy show a message. The payload's ``serial`` only goes
- up; a feed older than one already accepted is refused, so an old signed file
- cannot be re-served to bring back a withdrawn message.
- The feed is the full current list. On every accepted fetch the stored set is
- replaced: a message withdrawn upstream disappears here with its read markers.
- Any failure keeps the last good list.
- """
- from __future__ import annotations
- import asyncio
- import base64
- import json
- import logging
- import random
- import re
- from dataclasses import dataclass
- from datetime import datetime, timezone
- from urllib.parse import urlsplit
- import httpx
- from cryptography.exceptions import InvalidSignature
- from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey
- from sqlalchemy import delete, select
- from sqlalchemy.ext.asyncio import AsyncSession
- from backend.app.core.config import APP_VERSION
- from backend.app.models.announcement import Announcement, AnnouncementRead
- from backend.app.models.settings import Settings
- logger = logging.getLogger(__name__)
- FEED_URL = "https://raw.githubusercontent.com/maziggy/bambuddy-notifications/main/feed.json"
- FEED_FORMAT = 1
- # key_id (first 16 hex of sha256 of the raw public key) -> base64 raw public key.
- # A list so the key can be rotated: ship the new key next to the old one first.
- TRUSTED_KEYS: dict[str, str] = {
- "d70b3bf207fdfd59": "uQqUAYrInmXQ1nIwOIY/95L7tD5o/HZMmZfhAlULJ/w=",
- }
- # Where a message may link to. The registrar refuses anything else too; this is
- # the side that counts.
- LINK_HOSTS = ("github.com", "bambuddy.cool")
- LEVELS = ("info", "important", "critical")
- MAX_FEED_BYTES = 512 * 1024
- MAX_TITLE = 120
- MAX_BODY = 2000
- MAX_LINK_LABEL = 40
- FETCH_INTERVAL_SECONDS = 6 * 3600
- FETCH_JITTER_SECONDS = 30 * 60
- # Let startup settle before the first fetch.
- FIRST_FETCH_DELAY_SECONDS = 60
- ENABLED_KEY = "announcements_enabled"
- ALL_USERS_KEY = "announcements_all_users"
- # Internal state, never part of the settings API.
- SERIAL_KEY = "announcements_feed_serial"
- LAST_FETCH_KEY = "announcements_last_fetch"
- class FeedRejected(Exception):
- """The fetched file is not a feed this install accepts. Nothing changes."""
- # ---- verification -------------------------------------------------------------------
- def canonical(payload: dict) -> bytes:
- return json.dumps(payload, sort_keys=True, separators=(",", ":"), ensure_ascii=False).encode()
- def verify_feed(content: bytes, trusted_keys: dict[str, str] | None = None) -> dict:
- """The payload of a correctly signed feed, or FeedRejected."""
- keys = TRUSTED_KEYS if trusted_keys is None else trusted_keys
- try:
- envelope = json.loads(content)
- except (ValueError, UnicodeDecodeError) as exc:
- raise FeedRejected("not JSON") from exc
- if not isinstance(envelope, dict) or envelope.get("format") != FEED_FORMAT:
- raise FeedRejected("unknown feed format")
- public = keys.get(str(envelope.get("key_id")))
- if public is None:
- raise FeedRejected(f"signed with an unknown key ({envelope.get('key_id')!r})")
- payload = envelope.get("payload")
- if not isinstance(payload, dict):
- raise FeedRejected("no payload")
- try:
- signature = base64.b64decode(str(envelope.get("signature")), validate=True)
- Ed25519PublicKey.from_public_bytes(base64.b64decode(public)).verify(signature, canonical(payload))
- except (ValueError, InvalidSignature) as exc:
- raise FeedRejected("signature does not verify") from exc
- if payload.get("format") != FEED_FORMAT:
- raise FeedRejected("unknown payload format")
- serial = payload.get("serial")
- if not isinstance(serial, int) or isinstance(serial, bool) or serial < 1:
- raise FeedRejected("no serial")
- if not isinstance(payload.get("announcements"), list):
- raise FeedRejected("no announcement list")
- return payload
- # ---- what this install is ------------------------------------------------------------
- @dataclass(frozen=True)
- class InstallFacts:
- version: str
- channel: str # stable | beta
- install_type: str # docker | native | ha_addon | windows
- def _version_key(version: str) -> tuple:
- """Sortable form that puts a release above its own betas (1.2.6 > 1.2.6b3)."""
- from backend.app.api.routes.updates import parse_version
- parsed = parse_version(version)
- major, minor, patch, micro = (parsed + (0, 0, 0, 0))[:4]
- is_prerelease = parsed[4] if len(parsed) > 4 else 0
- prerelease_num = parsed[5] if len(parsed) > 5 else 0
- return (major, minor, patch, micro, 1 - is_prerelease, prerelease_num)
- def _install_type() -> str:
- from backend.app.api.routes import updates
- if updates._is_windows_installer_install():
- return "windows"
- if updates._is_ha_addon():
- return "ha_addon"
- if updates._is_docker_environment():
- return "docker"
- return "native"
- async def install_facts(db: AsyncSession) -> InstallFacts:
- beta_setting = (
- await db.execute(select(Settings.value).where(Settings.key == "include_beta_updates"))
- ).scalar_one_or_none()
- # Beta testers are the installs that asked for betas, and the ones running one.
- prerelease = bool(re.search(r"[a-zA-Z]", APP_VERSION.lstrip("v")))
- beta = prerelease or (beta_setting or "").lower() == "true"
- return InstallFacts(version=APP_VERSION, channel="beta" if beta else "stable", install_type=_install_type())
- def targets(entry: dict, facts: InstallFacts) -> bool:
- target = entry.get("target") or {}
- if not isinstance(target, dict):
- return False
- try:
- here = _version_key(facts.version)
- low, high = target.get("min_version"), target.get("max_version")
- if low and here < _version_key(str(low)):
- return False
- if high and here > _version_key(str(high)):
- return False
- except (TypeError, ValueError):
- return False
- channels = target.get("channels") or []
- if channels and facts.channel not in channels:
- return False
- install_types = target.get("install_types") or []
- return not (install_types and facts.install_type not in install_types)
- # ---- one entry, defensively ---------------------------------------------------------
- def link_allowed(url: str) -> bool:
- try:
- parts = urlsplit(url)
- port = parts.port
- except ValueError:
- return False
- host = (parts.hostname or "").lower()
- if parts.scheme != "https" or not host or parts.username or parts.password:
- return False
- if port not in (None, 443):
- return False
- return any(host == allowed or host.endswith("." + allowed) for allowed in LINK_HOSTS)
- def _parse_time(value: object) -> datetime | None:
- if not isinstance(value, str) or not value:
- return None
- try:
- parsed = datetime.fromisoformat(value.replace("Z", "+00:00"))
- except ValueError:
- return None
- if parsed.tzinfo is not None:
- parsed = parsed.astimezone(timezone.utc).replace(tzinfo=None)
- return parsed
- def _clean_texts(texts: object) -> dict[str, dict[str, str]] | None:
- """Plain strings, cut to length. None when there is no usable English text."""
- if not isinstance(texts, dict):
- return None
- cleaned: dict[str, dict[str, str]] = {}
- for lang, text in texts.items():
- if not isinstance(lang, str) or len(lang) > 10 or not isinstance(text, dict):
- continue
- title, body = text.get("title"), text.get("body")
- if not isinstance(title, str) or not isinstance(body, str) or not title.strip() or not body.strip():
- continue
- entry = {"title": title.strip()[:MAX_TITLE], "body": body.strip()[:MAX_BODY]}
- label = text.get("link_label")
- if isinstance(label, str) and label.strip():
- entry["link_label"] = label.strip()[:MAX_LINK_LABEL]
- cleaned[lang] = entry
- return cleaned if "en" in cleaned else None
- @dataclass
- class Entry:
- public_id: str
- level: str
- texts: dict[str, dict[str, str]]
- link_url: str | None
- published_at: datetime | None
- expires_at: datetime | None
- def parse_entry(raw: object) -> Entry | None:
- """One feed entry, or None if it can't be shown safely. Never raises."""
- if not isinstance(raw, dict):
- return None
- public_id = raw.get("id")
- if not isinstance(public_id, str) or not re.fullmatch(r"[A-Za-z0-9_-]{1,64}", public_id):
- return None
- level = raw.get("level") if raw.get("level") in LEVELS else "info"
- texts = _clean_texts(raw.get("texts"))
- if texts is None:
- return None
- link = raw.get("link_url")
- link = link if isinstance(link, str) and link_allowed(link) else None
- return Entry(
- public_id=public_id,
- level=level,
- texts=texts,
- link_url=link,
- published_at=_parse_time(raw.get("published_at")),
- expires_at=_parse_time(raw.get("expires_at")),
- )
- # ---- store ---------------------------------------------------------------------------
- async def _get(db: AsyncSession, key: str) -> str | None:
- return (await db.execute(select(Settings.value).where(Settings.key == key))).scalar_one_or_none()
- async def _set(db: AsyncSession, key: str, value: str) -> None:
- from backend.app.core.db_dialect import upsert_setting
- await upsert_setting(db, Settings, key, value)
- async def is_enabled(db: AsyncSession) -> bool:
- return (await _get(db, ENABLED_KEY) or "true").lower() != "false"
- async def apply_payload(db: AsyncSession, payload: dict, facts: InstallFacts) -> int:
- """Replace the stored announcements with the ones in ``payload`` that apply here.
- Refuses (FeedRejected) a serial below the highest already accepted. The same
- serial again is the same feed fetched twice, and is fine. Returns how many are
- stored. The caller commits.
- """
- serial = payload["serial"]
- seen = int(await _get(db, SERIAL_KEY) or 0)
- if serial < seen:
- raise FeedRejected(f"serial {serial} is older than {seen}, already accepted")
- wanted: dict[str, Entry] = {}
- for raw in payload["announcements"]:
- entry = parse_entry(raw)
- if entry is not None and targets(raw, facts):
- wanted[entry.public_id] = entry
- existing = {a.public_id: a for a in (await db.execute(select(Announcement))).scalars().all()}
- gone = [a.id for pid, a in existing.items() if pid not in wanted]
- if gone:
- # Read markers first: SQLite enforces no foreign keys unless asked to, so
- # ON DELETE CASCADE alone would leave them behind.
- await db.execute(delete(AnnouncementRead).where(AnnouncementRead.announcement_id.in_(gone)))
- await db.execute(delete(Announcement).where(Announcement.id.in_(gone)))
- for pid, entry in wanted.items():
- row = existing.get(pid) or Announcement(public_id=pid)
- row.level = entry.level
- row.texts = json.dumps(entry.texts, ensure_ascii=False)
- row.link_url = entry.link_url
- row.published_at = entry.published_at
- row.expires_at = entry.expires_at
- if pid not in existing:
- db.add(row)
- await _set(db, SERIAL_KEY, str(max(serial, seen)))
- return len(wanted)
- async def _download() -> bytes:
- # No version and no install identity in the request: the generic agent says
- # what is asking, nothing about which install.
- headers = {"User-Agent": "Bambuddy-Announcements", "Cache-Control": "no-cache"}
- async with (
- httpx.AsyncClient(timeout=15, follow_redirects=False) as client,
- client.stream("GET", FEED_URL, headers=headers) as response,
- ):
- if response.status_code != 200:
- raise FeedRejected(f"HTTP {response.status_code}")
- chunks: list[bytes] = []
- size = 0
- async for chunk in response.aiter_bytes():
- size += len(chunk)
- if size > MAX_FEED_BYTES:
- raise FeedRejected("larger than the feed is allowed to be")
- chunks.append(chunk)
- return b"".join(chunks)
- async def refresh(db: AsyncSession) -> int | None:
- """Fetch, verify and store. Returns how many apply here, or None if nothing changed.
- Never raises for a network problem or a bad feed -- those are logged and the
- last good list stays. A switched-off install fetches nothing.
- """
- if not await is_enabled(db):
- return None
- try:
- content = await _download()
- payload = verify_feed(content)
- count = await apply_payload(db, payload, await install_facts(db))
- except FeedRejected as exc:
- await db.rollback()
- logger.warning("Announcements feed ignored: %s", exc)
- return None
- except httpx.HTTPError as exc:
- await db.rollback()
- # Offline and air-gapped installs land here every time; not worth a warning.
- logger.info("Announcements feed not fetched: %s", type(exc).__name__)
- return None
- await _set(db, LAST_FETCH_KEY, datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"))
- await db.commit()
- logger.info("Announcements feed #%s: %d for this install", payload["serial"], count)
- return count
- # ---- read state ----------------------------------------------------------------------
- def _visible_filter(now: datetime):
- return (Announcement.expires_at.is_(None)) | (Announcement.expires_at > now)
- async def list_for(db: AsyncSession, user_id: int | None) -> list[dict]:
- """Live announcements, newest first, each with whether this user has read it."""
- now = datetime.now(timezone.utc).replace(tzinfo=None)
- rows = (
- (
- await db.execute(
- select(Announcement)
- .where(_visible_filter(now))
- .order_by(Announcement.published_at.desc(), Announcement.id.desc())
- )
- )
- .scalars()
- .all()
- )
- reader = AnnouncementRead.user_id.is_(None) if user_id is None else AnnouncementRead.user_id == user_id
- read_ids = set((await db.execute(select(AnnouncementRead.announcement_id).where(reader))).scalars().all())
- result = []
- for a in rows:
- try:
- texts = json.loads(a.texts)
- except (TypeError, ValueError):
- continue
- result.append(
- {
- "id": a.public_id,
- "level": a.level,
- "texts": texts,
- "link_url": a.link_url,
- "published_at": a.published_at.isoformat() + "Z" if a.published_at else None,
- "expires_at": a.expires_at.isoformat() + "Z" if a.expires_at else None,
- "read": a.id in read_ids,
- }
- )
- return result
- async def mark_read(db: AsyncSession, public_id: str, user_id: int | None) -> bool:
- """Record that this user read it. False if there is no such announcement."""
- announcement = (
- await db.execute(select(Announcement).where(Announcement.public_id == public_id))
- ).scalar_one_or_none()
- if announcement is None:
- return False
- reader = AnnouncementRead.user_id.is_(None) if user_id is None else AnnouncementRead.user_id == user_id
- # Checked rather than left to the unique constraint: NULL user_ids never
- # collide in one, so the auth-off row would be duplicated on every click.
- already = (
- await db.execute(select(AnnouncementRead.id).where(AnnouncementRead.announcement_id == announcement.id, reader))
- ).first()
- if already is None:
- db.add(AnnouncementRead(announcement_id=announcement.id, user_id=user_id))
- return True
- # ---- background loop -----------------------------------------------------------------
- _task: asyncio.Task | None = None
- async def _loop() -> None:
- from backend.app.core.database import async_session
- await asyncio.sleep(FIRST_FETCH_DELAY_SECONDS)
- while True:
- try:
- async with async_session() as db:
- await refresh(db)
- except asyncio.CancelledError:
- raise
- except Exception:
- logger.exception("Announcements refresh failed")
- await asyncio.sleep(FETCH_INTERVAL_SECONDS + random.uniform(0, FETCH_JITTER_SECONDS))
- def start() -> None:
- global _task
- if _task is None:
- _task = asyncio.create_task(_loop())
- def stop() -> None:
- global _task
- if _task is not None:
- _task.cancel()
- _task = None
|