From d67476ebbf73ab383e9f6db2ba430c88b4bc20c9 Mon Sep 17 00:00:00 2001 From: emi Date: Thu, 23 Jul 2026 04:41:19 +0000 Subject: [PATCH 1/3] notifier: bound push/Discord sends, cap pass duration, fix key-gen race (#5) --- indexnow.py | 64 ++++++++++++++--- notifications/discord.py | 12 +++- notifications/notify.py | 52 ++++++++++++-- notifications/push.py | 75 +++++++++++++++++--- tests/notifications/test_discord.py | 12 ++-- tests/notifications/test_discord_dm.py | 6 +- tests/notifications/test_notify.py | 71 +++++++++++++++++++ tests/notifications/test_push.py | 96 ++++++++++++++++++++++++++ tests/test_indexnow.py | 52 +++++++++++++- 9 files changed, 405 insertions(+), 35 deletions(-) create mode 100644 tests/notifications/test_push.py diff --git a/indexnow.py b/indexnow.py index 84bbc40..1bf3422 100644 --- a/indexnow.py +++ b/indexnow.py @@ -36,10 +36,58 @@ _STATE_PATH = os.environ.get("THERMOGRAPH_INDEXNOW_STATE_FILE") or os.path.join( _ENDPOINT = "https://api.indexnow.org/indexnow" # a shared endpoint; it fans out to all participants _BATCH = 10000 # IndexNow's max URLs per request +# One shared client instead of a bare httpx.post per batch — same reused-client +# pattern as web/app.py's _frontend_client and notifications/discord.py's _client. +# 30s default matches submit()'s own prior per-call default; still overridable +# per call (submit_all's fan-out passes its own `timeout` through unchanged). +_client = httpx.Client(timeout=30) + _lock = threading.Lock() _key = None +def _claim_file(value: str) -> str: + """Persist a freshly generated key as `_KEY_PATH`'s content, atomic + first-writer-wins — same race, same fix as push.py's VAPID `_claim_file` + (see its docstring): on a cold /state volume several workers can hit the + missing-file branch at once, and without this each would generate and cache + its OWN key, diverging from the file and from each other (the key served at + /{key}.txt would then randomly mismatch whichever key a given worker signed + a submission with). Write to a private temp file, then claim the real path + with a hard link (atomic; a loser gets EEXIST with no half-written window); + a loser reads back the winner's file instead of keeping its own generation.""" + tmp_path = f"{_KEY_PATH}.{os.getpid()}.tmp" + try: + os.makedirs(_DATA_DIR, exist_ok=True) + with open(tmp_path, "w", encoding="utf-8") as f: + f.write(value) + os.chmod(tmp_path, 0o600) + except OSError: + log.warning("could not persist IndexNow key to %s; using an in-memory key", _KEY_PATH) + return value + try: + try: + os.link(tmp_path, _KEY_PATH) + return value + except FileExistsError: + try: + with open(_KEY_PATH, encoding="utf-8") as f: + winner = f.read().strip() + if winner: + return winner + except OSError: + pass + return value + except OSError: + log.warning("could not persist IndexNow key to %s; using an in-memory key", _KEY_PATH) + return value + finally: + try: + os.unlink(tmp_path) + except OSError: + pass + + def key() -> str: """The IndexNow key (env → file → generate-and-persist), cached for the process.""" global _key @@ -60,14 +108,10 @@ def key() -> str: return _key except OSError: pass - _key = secrets.token_hex(16) # 32 hex chars — within IndexNow's 8–128 range - try: - os.makedirs(_DATA_DIR, exist_ok=True) - with open(_KEY_PATH, "w", encoding="utf-8") as f: - f.write(_key) - os.chmod(_KEY_PATH, 0o600) - except OSError: - log.warning("could not persist IndexNow key to %s; using an in-memory key", _KEY_PATH) + # Cache whatever _claim_file settles on — see its docstring: our own + # generation if we won the race to create the file, the winner's key + # read back from disk if we lost it. + _key = _claim_file(secrets.token_hex(16)) # 32 hex chars — within IndexNow's 8–128 range return _key @@ -83,8 +127,8 @@ def submit(urls, host: str, key_location: str, scheme: str = "https", timeout: f batch = urls[i:i + _BATCH] payload = {"host": host, "key": k, "keyLocation": key_location, "urlList": batch} try: - r = httpx.post(_ENDPOINT, json=payload, timeout=timeout, - headers={"Content-Type": "application/json; charset=utf-8"}) + r = _client.post(_ENDPOINT, json=payload, timeout=timeout, + headers={"Content-Type": "application/json; charset=utf-8"}) except httpx.HTTPError as e: log.warning("IndexNow request failed: %s", e) statuses.append("error") diff --git a/notifications/discord.py b/notifications/discord.py index 776fa2f..2eb0e24 100644 --- a/notifications/discord.py +++ b/notifications/discord.py @@ -47,6 +47,12 @@ _COLD = 0x4393C3 _TIMEOUT = httpx.Timeout(10.0) +# One shared client instead of a bare httpx.post per call — same reused-client +# pattern as web/app.py's _frontend_client, so every webhook/bot-REST call in +# this module reuses one connection pool instead of paying a fresh TCP+TLS +# handshake per notification. +_client = httpx.Client(timeout=_TIMEOUT) + def enabled() -> bool: return bool(WEBHOOK_URL) or weather_enabled() @@ -131,7 +137,7 @@ def post_daily_feed(feed: dict | None = None) -> bool: "embeds": [embed], } try: - resp = httpx.post(WEBHOOK_URL, json=payload, timeout=_TIMEOUT) + resp = _client.post(WEBHOOK_URL, json=payload, timeout=_TIMEOUT) ok = 200 <= resp.status_code < 300 except Exception: # noqa: BLE001 - best-effort side channel pass @@ -159,10 +165,10 @@ def _bot_post(path: str, json_body: dict) -> httpx.Response | None: """POST to the bot REST API with one 429 retry. None on a transport error.""" headers = {"Authorization": f"Bot {BOT_TOKEN}"} try: - resp = httpx.post(f"{_API}{path}", json=json_body, headers=headers, timeout=_TIMEOUT) + resp = _client.post(f"{_API}{path}", json=json_body, headers=headers, timeout=_TIMEOUT) if resp.status_code == 429: time.sleep(_retry_after(resp)) - resp = httpx.post(f"{_API}{path}", json=json_body, headers=headers, timeout=_TIMEOUT) + resp = _client.post(f"{_API}{path}", json=json_body, headers=headers, timeout=_TIMEOUT) return resp except Exception: # noqa: BLE001 - best-effort side channel return None diff --git a/notifications/notify.py b/notifications/notify.py index d7c82d9..a26702f 100644 --- a/notifications/notify.py +++ b/notifications/notify.py @@ -23,6 +23,7 @@ A forecast-fetch failure just skips the cell, so a pass never dies on a rate lim """ import concurrent.futures import datetime +import logging import os import threading import time @@ -42,6 +43,8 @@ from notifications import push from accounts.db import sync_session_maker from accounts.models import AccessToken, Notification, PushSubscription, Subscription, User +log = logging.getLogger("thermograph.notify") + _BASE = os.environ.get("THERMOGRAPH_BASE", "/thermograph").strip("/") BASE = f"/{_BASE}" if _BASE else "" @@ -57,6 +60,14 @@ FORECAST_HORIZON_DAYS = 7 # — the rest are picked up on later passes. MAX_ARCHIVE_FETCHES_PER_PASS = int(os.environ.get("THERMOGRAPH_NOTIFY_ARCHIVE_FETCHES", "4")) +# Soft wall-clock budget for one whole pass, same discipline as the archive-fetch +# cap above: checked between cells (never mid-cell) so one pathological pass — a +# slow upstream, a DB hiccup, hundreds of cells each doing real work — can't run +# unbounded and push the next pass's wake past its own interval. 2x INTERVAL by +# default: generous enough that an ordinarily-slow pass still finishes, but a +# genuinely runaway one still yields before it would overlap the next wake. +PASS_DEADLINE_SECONDS = int(os.environ.get("THERMOGRAPH_NOTIFY_PASS_DEADLINE", str(INTERVAL * 2))) + # Metrics whose low tail also counts as unusual when a subscription is two-sided # (cold snaps, unusually calm/dry). Precipitation is one-directional (high only). TWO_SIDED_METRICS = set(grading.TEMP_METRICS) @@ -171,19 +182,40 @@ def _send_discord_job(job: tuple[str, str, str, str]) -> None: # hundreds of subscribers no longer stretches a pass out send-by-send. SEND_WORKERS = int(os.environ.get("THERMOGRAPH_NOTIFY_SEND_WORKERS", "8")) +# Belt-and-suspenders on top of push.py's own send timeout: f.result() below has +# no timeout by default and would otherwise wait forever for a future that never +# resolves, which — same as an untimed webpush() call — wedges the notifier for +# the process's life under the singleton flock. Generous on purpose; this only +# needs to fire if a per-send timeout was somehow bypassed, not to race it. +SEND_RESULT_TIMEOUT = int(os.environ.get("THERMOGRAPH_NOTIFY_SEND_RESULT_TIMEOUT", "60")) + def _flush_sends(push_jobs, discord_jobs) -> list[int]: """Deliver every gathered push/Discord job concurrently. Returns the push-subscription row ids the push service reported gone, for the caller - to prune. Blocks until every job has finished (or errored) — a pass still - waits for delivery to complete, it just no longer does so serially.""" + to prune. Blocks until every job has finished or errored, and gives up + waiting on (but does not lose track of) any single job past + SEND_RESULT_TIMEOUT — a pass still waits for delivery to complete, it just + no longer does so serially, and no misbehaving future can hang the result + loop past a bounded wait.""" with concurrent.futures.ThreadPoolExecutor(max_workers=SEND_WORKERS) as pool: futures = [pool.submit(_send_push_job, job) for job in push_jobs] futures += [pool.submit(_send_discord_job, job) for job in discord_jobs] gone = [] for f in futures: try: - result = f.result() + result = f.result(timeout=SEND_RESULT_TIMEOUT) + except concurrent.futures.TimeoutError: + # f.result() gave up waiting, but the worker thread itself can't be + # killed — ThreadPoolExecutor.__exit__ still calls shutdown(wait=True), + # so the pool as a whole won't return until that thread actually + # finishes. The real backstop is the per-call socket timeouts this + # commit adds (push.py's _SEND_TIMEOUT, discord.py's _TIMEOUT): they + # bound how long a "stuck" send can truly run. This catch just stops + # THIS result from being reported/counted as gone, and keeps checking + # the remaining futures instead of getting stuck on one. + log.warning("push/Discord send exceeded %ss timeout; treating as failed", SEND_RESULT_TIMEOUT) + continue except Exception: # noqa: BLE001 - one bad send must not lose the rest continue if result is not None: @@ -370,7 +402,18 @@ def run_pass() -> int: by_cell: dict[str, list] = {} for sub in subs: by_cell.setdefault(sub.cell_id, []).append(sub) - for cell_id, cell_subs in by_cell.items(): + cells_skipped = 0 + for i, (cell_id, cell_subs) in enumerate(by_cell.items()): + # Checked BETWEEN cells, not inside _process_cell — a soft budget, not + # a hard per-item timeout. Once tripped, every remaining cell this pass + # is skipped outright (not attempted-and-abandoned mid-flight) so a + # skipped cell's subscriptions are simply picked up whole on the next + # pass, same as a rate-limited one already is. + if time.perf_counter() - t0 > PASS_DEADLINE_SECONDS: + cells_skipped = len(by_cell) - i + log.warning("notifier pass exceeded %ss deadline; skipping %d/%d remaining cell(s)", + PASS_DEADLINE_SECONDS, cells_skipped, len(by_cell)) + break try: created += _process_cell(session, cell_id, cell_subs, today, now, archive_budget) except Exception: # noqa: BLE001 - one bad cell must not abort the pass @@ -382,6 +425,7 @@ def run_pass() -> int: audit.log_activity("notify.pass", { "subs_evaluated": len(subs), "cells": len(by_cell), "created": created, "archives_fetched": MAX_ARCHIVE_FETCHES_PER_PASS - archive_budget[0], + "cells_skipped_deadline": cells_skipped, "duration_ms": round((time.perf_counter() - t0) * 1000.0, 1)}) return created diff --git a/notifications/push.py b/notifications/push.py index 0c40a51..594607f 100644 --- a/notifications/push.py +++ b/notifications/push.py @@ -39,6 +39,15 @@ _VAPID_PATH = os.environ.get("THERMOGRAPH_VAPID_FILE") or os.path.join(_DATA_DIR # The VAPID "sub" claim — a contact the push service can reach about our traffic. _CONTACT = os.environ.get("THERMOGRAPH_VAPID_CONTACT", "mailto:admin@thermograph.app") +# pywebpush 2.0.0 forwards `timeout` straight to requests.post with no default of +# its own — omit the kwarg here and the send blocks with NO timeout at all (the +# "10s default" people expect only applies if requests itself is called bare). +# One hung push endpoint must never wedge the notifier: the singleton flock +# (core/singleton.py) is held for the process's whole life, so no other worker +# can take over while this thread is stuck, and /healthz keeps reporting green +# on a notifier that's actually frozen. +_SEND_TIMEOUT = 15 + _lock = threading.Lock() _keys = None # cached {"private_key": str, "public_key": str} (base64url-raw) @@ -60,6 +69,59 @@ def _generate() -> dict: return {"private_key": _b64url(scalar), "public_key": _b64url(raw_pub)} +def _claim_file(data: dict) -> dict: + """Persist a freshly generated keypair as `_VAPID_PATH`'s content, atomic + first-writer-wins. `_lock` only keeps this process's own threads from racing + each other — on a cold /state volume every uvicorn *worker* (a separate + process) hits the missing-file branch at boot together, and without this each + would generate its OWN keypair and cache it in-process, silently diverging + from the file and from each other (a subscription signed against one worker's + public key fails to verify under another's private key — the exact incident + class deploy/entrypoint.sh's comments name). + + Write the full keypair to a private temp file first, then claim the real path + with a hard link: `os.link` is atomic and the loser gets EEXIST immediately, + with no window where `_VAPID_PATH` exists but is only half-written (unlike + O_CREAT|O_EXCL directly on the destination followed by a separate write). A + loser discards its own generation and reads back the winner's file instead, + so every process ends up caching the SAME keys the file actually holds.""" + tmp_path = f"{_VAPID_PATH}.{os.getpid()}.tmp" + try: + os.makedirs(_DATA_DIR, exist_ok=True) + with open(tmp_path, "w", encoding="utf-8") as f: + json.dump(data, f) + os.chmod(tmp_path, 0o600) + except OSError: + log.warning("could not persist VAPID keys to %s; using an in-memory pair", _VAPID_PATH) + return data + try: + try: + os.link(tmp_path, _VAPID_PATH) + return data + except FileExistsError: + # Lost the race — someone else's file is now the truth. Read IT back + # rather than keep our own now-orphaned generation. + try: + with open(_VAPID_PATH, encoding="utf-8") as f: + winner = json.load(f) + if winner.get("private_key") and winner.get("public_key"): + return winner + except (OSError, ValueError): + pass + # Winner's file was unreadable/corrupt (very rare) — fall back to our + # own in-memory generation rather than crash; the next _load() call + # (a new process, or after the file heals) retries the file. + return data + except OSError: + log.warning("could not persist VAPID keys to %s; using an in-memory pair", _VAPID_PATH) + return data + finally: + try: + os.unlink(tmp_path) + except OSError: + pass + + def _load() -> dict: """Resolve the keypair once (env → file → generate) and cache it.""" global _keys @@ -81,14 +143,10 @@ def _load() -> dict: return _keys except (OSError, ValueError): pass - _keys = _generate() - try: - os.makedirs(_DATA_DIR, exist_ok=True) - with open(_VAPID_PATH, "w", encoding="utf-8") as f: - json.dump(_keys, f) - os.chmod(_VAPID_PATH, 0o600) - except OSError: - log.warning("could not persist VAPID keys to %s; using an in-memory pair", _VAPID_PATH) + # Cache whatever _claim_file settles on — our own generation if we won + # the race to create the file, or the winner's keys read back from disk + # if we lost it. Either way this process's cache matches the file. + _keys = _claim_file(_generate()) return _keys @@ -111,6 +169,7 @@ def send(subscription_info: dict, payload: dict) -> str: vapid_private_key=keys["private_key"], vapid_claims={"sub": _CONTACT}, ttl=86400, + timeout=_SEND_TIMEOUT, ) return "ok" except WebPushException as e: diff --git a/tests/notifications/test_discord.py b/tests/notifications/test_discord.py index f62763e..64a50a6 100644 --- a/tests/notifications/test_discord.py +++ b/tests/notifications/test_discord.py @@ -1,5 +1,5 @@ """Discord daily-feed post: embed shape, webhook delivery, and the once-a-day guard. -No network — httpx.post is stubbed, like the IndexNow tests.""" +No network — the shared client's .post is stubbed, like the IndexNow tests.""" import datetime import time @@ -63,7 +63,7 @@ def test_build_embed_rejects_empty_feed(): def test_post_disabled_without_webhook(monkeypatch): monkeypatch.setattr(discord, "WEBHOOK_URL", "") posted = [] - monkeypatch.setattr(discord.httpx, "post", lambda *a, **k: posted.append(a) or None) + monkeypatch.setattr(discord._client, "post", lambda *a, **k: posted.append(a) or None) assert discord.post_daily_feed(_feed([_CARD_HOT])) is False assert posted == [] # never touched the network @@ -75,7 +75,7 @@ class _Resp: def test_post_sends_embed_to_the_webhook(monkeypatch): monkeypatch.setattr(discord, "WEBHOOK_URL", "https://discord.test/webhook/abc") calls = [] - monkeypatch.setattr(discord.httpx, "post", + monkeypatch.setattr(discord._client, "post", lambda url, json=None, **k: (calls.append((url, json)), _Resp(204))[1]) assert discord.post_daily_feed(_feed([_CARD_HOT])) is True (url, payload), = calls @@ -91,7 +91,7 @@ def test_post_also_broadcasts_to_weather_channel(monkeypatch): monkeypatch.setattr(discord, "BOT_TOKEN", "bot-abc") monkeypatch.setattr(discord, "WEATHER_CHANNEL_ID", "chan-weather") calls = [] - monkeypatch.setattr(discord.httpx, "post", + monkeypatch.setattr(discord._client, "post", lambda url, json=None, **k: (calls.append((url, json)), _Resp(204))[1]) assert discord.post_daily_feed(_feed([_CARD_HOT])) is True (url, body), = calls @@ -104,14 +104,14 @@ def test_post_is_best_effort_on_http_error(monkeypatch): def _boom(*a, **k): raise RuntimeError("network down") - monkeypatch.setattr(discord.httpx, "post", _boom) + monkeypatch.setattr(discord._client, "post", _boom) # Never raises; just reports failure. assert discord.post_daily_feed(_feed([_CARD_HOT])) is False def test_post_skips_stale_or_empty_feed(monkeypatch): monkeypatch.setattr(discord, "WEBHOOK_URL", "https://discord.test/webhook/abc") - monkeypatch.setattr(discord.httpx, "post", lambda *a, **k: _Resp(204)) + monkeypatch.setattr(discord._client, "post", lambda *a, **k: _Resp(204)) # Yesterday's feed is stale -> no post. yesterday = (datetime.date.today() - datetime.timedelta(days=1)).isoformat() assert discord.post_daily_feed(_feed([_CARD_HOT], date=yesterday)) is False diff --git a/tests/notifications/test_discord_dm.py b/tests/notifications/test_discord_dm.py index 1e1b34e..aa9cc56 100644 --- a/tests/notifications/test_discord_dm.py +++ b/tests/notifications/test_discord_dm.py @@ -28,7 +28,7 @@ class _Resp: def _mock_posts(monkeypatch, script): - """Route discord.httpx.post by URL; `script` maps a URL substring -> _Resp.""" + """Route discord._client.post by URL; `script` maps a URL substring -> _Resp.""" calls = [] def _post(url, json=None, headers=None, timeout=None): @@ -37,7 +37,7 @@ def _mock_posts(monkeypatch, script): if frag in url: return resp return _Resp(404) - monkeypatch.setattr(discord.httpx, "post", _post) + monkeypatch.setattr(discord._client, "post", _post) return calls @@ -97,7 +97,7 @@ def test_bot_post_retries_once_on_429(monkeypatch): monkeypatch.setattr(discord, "BOT_TOKEN", "bot-abc") monkeypatch.setattr(discord.time, "sleep", lambda s: None) # don't actually wait seq = iter([_Resp(429, {"retry_after": 0.01}), _Resp(200, {"id": "c"})]) - monkeypatch.setattr(discord.httpx, "post", lambda *a, **k: next(seq)) + monkeypatch.setattr(discord._client, "post", lambda *a, **k: next(seq)) resp = discord._bot_post("/users/@me/channels", {"recipient_id": "d"}) assert resp.status_code == 200 # the retry succeeded diff --git a/tests/notifications/test_notify.py b/tests/notifications/test_notify.py index 8c22038..74399f1 100644 --- a/tests/notifications/test_notify.py +++ b/tests/notifications/test_notify.py @@ -2,12 +2,14 @@ wording (pure logic), plus an integration check that a full pass fetches a MISSING archive once but never re-fetches a cached one (against conftest's throwaway DB).""" import datetime +import time import types import uuid import numpy as np import polars as pl +from core import audit from data import climate from accounts import db from notifications import notify @@ -211,3 +213,72 @@ def test_inapp_notification_survives_push_error(monkeypatch): assert len(_user_notifications(uid)) == 1 # the in-app write is unaffected assert _push_count() == 1 # a mere error doesn't prune + + +# --- send-result timeout: one hung future must not wedge the pass ----------- +def test_flush_sends_treats_timeout_as_a_failed_send(monkeypatch): + monkeypatch.setattr(notify, "SEND_RESULT_TIMEOUT", 0.05) + + def _hang(job): + time.sleep(0.3) # longer than the timeout above + return None + monkeypatch.setattr(notify, "_send_push_job", _hang) + + # Must not raise, and a timed-out send is never reported as a "gone" + # endpoint to prune (it never actually got a definitive answer). + gone = notify._flush_sends([object()], []) + assert gone == [] + + +def test_flush_sends_still_returns_gone_ids_within_the_timeout(monkeypatch): + monkeypatch.setattr(notify, "SEND_RESULT_TIMEOUT", 5) + monkeypatch.setattr(notify, "_send_push_job", lambda job: job) # echoes the id straight back + gone = notify._flush_sends([7, 8], []) + assert sorted(gone) == [7, 8] + + +# --- pass wall-clock deadline: one pathological pass can't run unbounded ---- +def _seed_two_cell_subscriptions(): + db.Base.metadata.create_all(db.sync_engine) + with db.sync_session_maker() as s: + s.execute(delete(User)) + s.commit() + uid1, uid2 = uuid.uuid4(), uuid.uuid4() + s.add(User(id=uid1, email="deadline1@example.com", hashed_password="x", is_active=True)) + s.add(User(id=uid2, email="deadline2@example.com", hashed_password="x", is_active=True)) + s.commit() + s.add(Subscription(user_id=uid1, cell_id="500_600", label="A", lat=1.0, lon=2.0, + threshold=95, metrics=["tmax"], kind="observed", two_sided=False)) + s.add(Subscription(user_id=uid2, cell_id="700_800", label="B", lat=3.0, lon=4.0, + threshold=95, metrics=["tmax"], kind="observed", two_sided=False)) + s.commit() + return uid1, uid2 + + +def test_run_pass_stops_at_wall_clock_deadline(monkeypatch): + uid1, uid2 = _seed_two_cell_subscriptions() + _cached_extreme(monkeypatch) # every cell would otherwise trigger a notification + monkeypatch.setattr(notify, "PASS_DEADLINE_SECONDS", 0) # trips before the first cell + + logged = [] + monkeypatch.setattr(audit, "log_activity", + lambda kind, data: logged.append((kind, data))) + + notify.run_pass() + + # A deadline of 0 must skip every cell rather than process any of them. + assert _user_notifications(uid1) == [] + assert _user_notifications(uid2) == [] + pass_log = next(data for kind, data in logged if kind == "notify.pass") + assert pass_log["cells_skipped_deadline"] == 2 + + +def test_run_pass_processes_normally_within_the_deadline(monkeypatch): + uid1, uid2 = _seed_two_cell_subscriptions() + _cached_extreme(monkeypatch) + monkeypatch.setattr(notify, "PASS_DEADLINE_SECONDS", 300) # generous — shouldn't trip + + notify.run_pass() + + assert len(_user_notifications(uid1)) == 1 + assert len(_user_notifications(uid2)) == 1 diff --git a/tests/notifications/test_push.py b/tests/notifications/test_push.py new file mode 100644 index 0000000..422583f --- /dev/null +++ b/tests/notifications/test_push.py @@ -0,0 +1,96 @@ +"""VAPID key resolution: the atomic first-writer-wins claim of the keypair file +(push.py's `_claim_file`), and that `_load()` caches whatever that settles on +rather than its own local generation when it loses the race. + +No network — `send()`'s pywebpush call isn't exercised here (that's notify.py's +`test_push_dispatched_on_new_notification` etc., which stub `push.send` itself).""" +import json + +from notifications import push + + +def _reset_state(monkeypatch, path): + monkeypatch.setattr(push, "_VAPID_PATH", str(path)) + monkeypatch.setattr(push, "_DATA_DIR", str(path.parent)) + monkeypatch.setattr(push, "_keys", None) + monkeypatch.delenv("THERMOGRAPH_VAPID_PRIVATE_KEY", raising=False) + monkeypatch.delenv("THERMOGRAPH_VAPID_PUBLIC_KEY", raising=False) + + +def test_claim_file_first_writer_wins(monkeypatch, tmp_path): + path = tmp_path / "vapid.json" + _reset_state(monkeypatch, path) + data = {"private_key": "priv-a", "public_key": "pub-a"} + + result = push._claim_file(data) + + assert result == data + assert json.loads(path.read_text()) == data + # No leftover temp file. + assert list(tmp_path.iterdir()) == [path] + + +def test_claim_file_reads_back_winner_on_race(monkeypatch, tmp_path): + """A process that loses the os.link() race must discard its own freshly + generated keypair and use whatever the winner actually persisted.""" + path = tmp_path / "vapid.json" + _reset_state(monkeypatch, path) + winner = {"private_key": "priv-winner", "public_key": "pub-winner"} + path.write_text(json.dumps(winner)) # simulates another process winning first + loser = {"private_key": "priv-loser", "public_key": "pub-loser"} + + result = push._claim_file(loser) + + assert result == winner + assert json.loads(path.read_text()) == winner # the loser never touched the file + + +def test_load_caches_winner_keys_not_its_own_generation(monkeypatch, tmp_path): + """End-to-end through _load(): on a cold /state volume, a worker that hits the + generate branch but loses the write race must end up with the SAME in-process + cached keys the file actually holds — not the keys it generated locally, + which would silently diverge from every other worker (a subscription signed + against one worker's public key then fails to verify under another's + private key).""" + path = tmp_path / "vapid.json" + _reset_state(monkeypatch, path) + winner = {"private_key": "priv-winner", "public_key": "pub-winner"} + + def _generate_and_lose_the_race(): + # Between our cache-miss file read and our own _claim_file() call, a + # concurrent worker wins and writes the real file first. + path.write_text(json.dumps(winner)) + return {"private_key": "priv-mine", "public_key": "pub-mine"} + + monkeypatch.setattr(push, "_generate", _generate_and_lose_the_race) + + result = push._load() + + assert result == winner + assert push._keys == winner # the process-wide cache matches the file + assert push.public_key() == "pub-winner" + + +def test_load_reads_existing_file_without_generating(monkeypatch, tmp_path): + path = tmp_path / "vapid.json" + _reset_state(monkeypatch, path) + existing = {"private_key": "priv-x", "public_key": "pub-x"} + path.write_text(json.dumps(existing)) + + def boom(): + raise AssertionError("must not generate when a valid file already exists") + monkeypatch.setattr(push, "_generate", boom) + + assert push._load() == existing + + +def test_load_prefers_env_over_file(monkeypatch, tmp_path): + path = tmp_path / "vapid.json" + _reset_state(monkeypatch, path) + path.write_text(json.dumps({"private_key": "priv-file", "public_key": "pub-file"})) + monkeypatch.setenv("THERMOGRAPH_VAPID_PRIVATE_KEY", " priv-env ") + monkeypatch.setenv("THERMOGRAPH_VAPID_PUBLIC_KEY", " pub-env ") + + result = push._load() + + assert result == {"private_key": "priv-env", "public_key": "pub-env"} diff --git a/tests/test_indexnow.py b/tests/test_indexnow.py index d8a7c42..9045ad1 100644 --- a/tests/test_indexnow.py +++ b/tests/test_indexnow.py @@ -69,7 +69,7 @@ def test_indexnow_submit_all_builds_payload(monkeypatch): status_code = 200 text = "ok" - monkeypatch.setattr(indexnow.httpx, "post", + monkeypatch.setattr(indexnow._client, "post", lambda url, json=None, **kw: (calls.append((url, json)), _Resp())[1]) res = indexnow.submit_all("https://thermograph.org") assert res["submitted"] == res["total"] > 100 @@ -87,3 +87,53 @@ def test_indexnow_if_changed_state(monkeypatch, tmp_path): assert indexnow._read_state() == "" # first deploy → would submit indexnow._write_state(sig) assert indexnow._read_state() == sig # unchanged → deploy would skip + + +# --- key() race: atomic first-writer-wins on a cold /state volume ----------- +# Same race, same fix as push.py's VAPID keypair (see tests/notifications/ +# test_push.py) — several workers can hit the missing-key-file branch at once; +# only one may actually create the file, everyone else must read IT back. + +def _reset_key_state(monkeypatch, path): + monkeypatch.setattr(indexnow, "_KEY_PATH", str(path)) + monkeypatch.setattr(indexnow, "_DATA_DIR", str(path.parent)) + monkeypatch.setattr(indexnow, "_key", None) + monkeypatch.delenv("THERMOGRAPH_INDEXNOW_KEY", raising=False) + + +def test_claim_file_first_writer_wins(monkeypatch, tmp_path): + path = tmp_path / "indexnow_key.txt" + _reset_key_state(monkeypatch, path) + + result = indexnow._claim_file("key-a") + + assert result == "key-a" + assert path.read_text() == "key-a" + assert list(tmp_path.iterdir()) == [path] # no leftover temp file + + +def test_claim_file_reads_back_winner_on_race(monkeypatch, tmp_path): + path = tmp_path / "indexnow_key.txt" + _reset_key_state(monkeypatch, path) + path.write_text("key-winner") # simulates another process winning first + + result = indexnow._claim_file("key-loser") + + assert result == "key-winner" + assert path.read_text() == "key-winner" # the loser never touched the file + + +def test_key_caches_winner_not_its_own_generation(monkeypatch, tmp_path): + path = tmp_path / "indexnow_key.txt" + _reset_key_state(monkeypatch, path) + + def _lose_the_race(*a, **k): + path.write_text("key-winner") + return "a" * 32 # what secrets.token_hex(16) would have produced locally + + monkeypatch.setattr(indexnow.secrets, "token_hex", _lose_the_race) + + result = indexnow.key() + + assert result == "key-winner" + assert indexnow._key == "key-winner" # the process-wide cache matches the file From 215c46cea75eaa0546c875bdc854f31823ec3ae2 Mon Sep 17 00:00:00 2001 From: emi Date: Thu, 23 Jul 2026 04:41:26 +0000 Subject: [PATCH 2/3] Bound Postgres connections, add rate-limit/timeout guards, move revgeo off the threadpool (#8) --- accounts/db.py | 32 +++- data/climate.py | 292 ++++++++++++++++++++++-------- data/climate_store.py | 215 +++++++++++++--------- data/places.py | 9 +- data/store.py | 194 +++++++++++++------- requirements.txt | 6 +- tests/accounts/test_db_engines.py | 25 +++ tests/data/test_climate.py | 197 ++++++++++++++++++++ tests/data/test_climate_store.py | 11 +- tests/data/test_store.py | 88 +++++++++ tests/test_warm_cities.py | 38 ++++ tests/web/test_gzip.py | 12 ++ warm_cities.py | 21 +++ web/app.py | 11 +- 14 files changed, 912 insertions(+), 239 deletions(-) create mode 100644 tests/test_warm_cities.py create mode 100644 tests/web/test_gzip.py diff --git a/accounts/db.py b/accounts/db.py index 680c15f..8b97ff8 100644 --- a/accounts/db.py +++ b/accounts/db.py @@ -56,18 +56,40 @@ def _apply_sqlite_pragmas(dbapi_conn, _rec): cur.close() +# Postgres-only pool sizing + sync-path timeouts, defined at module scope (not +# only inside `if IS_POSTGRES` below) so the values themselves are unit-testable +# without a real Postgres connection. +# +# pool_size=1/max_overflow=1 (the original setting on both) meant a 3rd +# concurrent request per worker had to wait out pool_timeout (~30s default) for a +# slot — easy to hit with 4 uvicorn workers x real traffic. 5+5 gives each async +# engine headroom for a burst of concurrent account requests without coming close +# to Postgres's own connection ceiling (two async engines here plus the sync +# engine below, times N workers, all share that budget). +_ASYNC_POOL_KWARGS = dict(pool_size=5, max_overflow=5, pool_pre_ping=True, + pool_recycle=1800, future=True) +# Smaller than the async engines: only the notifier leader (one process +# cluster-wide, see core/singleton.py) uses the sync engine, driving its periodic +# sweep from a plain thread with no event loop. connect_timeout + a +# statement_timeout (via psycopg's `options`) bound how long a wedged Postgres +# can hang the notifier — it runs for the lifetime of that process's leader +# flock, so an indefinite hang here would silently stop every subscription +# notification until the process is restarted. +_SYNC_POOL_KWARGS = dict(pool_size=2, max_overflow=3, pool_pre_ping=True, + pool_recycle=1800, future=True, + connect_args={"connect_timeout": 5, + "options": "-c statement_timeout=30000"}) + if IS_POSTGRES: DB_PATH = None - _COMMON = dict(pool_size=1, max_overflow=1, pool_pre_ping=True, pool_recycle=1800, future=True) # Read-write engine (default) + a read-only engine that pins every session to a # read-only transaction on the server, so a pure-read endpoint can't accidentally # write and a long read stays off the write connection. - async_engine = create_async_engine(DATABASE_URL, **_COMMON) + async_engine = create_async_engine(DATABASE_URL, **_ASYNC_POOL_KWARGS) async_engine_ro = create_async_engine( - DATABASE_URL, **_COMMON, + DATABASE_URL, **_ASYNC_POOL_KWARGS, connect_args={"server_settings": {"default_transaction_read_only": "on"}}) - sync_engine = create_engine(_sync_url(DATABASE_URL), pool_size=1, max_overflow=1, - pool_pre_ping=True, pool_recycle=1800, future=True) + sync_engine = create_engine(_sync_url(DATABASE_URL), **_SYNC_POOL_KWARGS) else: DB_PATH = os.environ.get("THERMOGRAPH_ACCOUNTS_DB") or os.path.join(paths.DATA_DIR, "accounts.sqlite") os.makedirs(os.path.dirname(DB_PATH), exist_ok=True) diff --git a/data/climate.py b/data/climate.py index f88fc43..8f22ce1 100644 --- a/data/climate.py +++ b/data/climate.py @@ -5,8 +5,11 @@ a zstd-compressed parquet file keyed by cell id. Percentiles are computed on the fly from this raw record (see grading.py), which keeps the cache small and lets us handle ties (e.g. many zero-precip days) correctly. """ +import concurrent.futures import datetime import os +import queue +import tempfile import threading import time @@ -20,8 +23,30 @@ import paths from data import climate_store from data import store +# Reused across every fetch (archive/forecast/geocode/revgeo) instead of a bare +# httpx.get per call opening a fresh TCP+TLS connection every time -- mirrors +# web/app.py's reused _frontend_client. Every call still passes its own explicit +# per-call timeout (see _request), so this only changes connection reuse, never +# retry/backoff semantics. Never explicitly closed: this is a long-lived module +# used by both the server process and one-shot scripts (warm_cities.py, +# migrate_cache_to_pg.py), and the OS reclaims the sockets at process exit either +# way. +_client = httpx.Client() + CACHE_DIR = os.path.join(paths.DATA_DIR, "cache") MAX_ATTEMPTS = 3 # per upstream call, before giving up + +# _fetch_history / _fetch_history_nasa run under _load_history's per-cell lock +# (see _cell_lock), so every attempt's timeout adds directly to how long a +# same-cell waiter queues behind this thread. A flat 180s x MAX_ATTEMPTS could +# pin the thread for up to 9 minutes on one slow-but-eventually-ok upstream. The +# full archive response is genuinely large (~16k daily rows from START_DATE, see +# below), so it can legitimately need more than a topup fetch's 60s -- but that +# headroom belongs on the LAST attempt only: the first attempts fail fast so a +# truly wedged upstream frees the lock for a waiter sooner, and only the final, +# worth-waiting-out attempt gets the longer allowance. Worst case under the lock +# drops from 540s (9 min) to 60+60+150=270s (4.5 min). +ARCHIVE_FETCH_TIMEOUTS = (60, 60, 150) START_DATE = "1980-01-01" # ERA5 reaches back to 1940; 1980 = 45 yrs, fast + robust ARCHIVE_LATENCY_DAYS = 6 # ERA5 archive lags real time by a few days # A real archive spans decades (~16k daily rows from START_DATE). A self-hosted @@ -111,6 +136,17 @@ _archive_cooldown_until = 0.0 _archive_limit_daily = False # is the current cooldown a daily-quota exhaustion? +# The forecast (recent+forward) endpoint's own cooldown, tracked separately from +# the archive one above: Open-Meteo's archive and forecast APIs have independent +# quotas, so a 429 on one must not gate (or be masked by) a cooldown meant for the +# other. Without this, a forecast brownout let every subscribed cell AND every +# live /cell request independently retry-storm the endpoint -- each paying +# MAX_ATTEMPTS x up to 60s plus the MET Norway fallback's own retry cycle, with no +# circuit breaker (see _load_recent_forecast). +FORECAST_COOLDOWN = 120 # seconds +_forecast_cooldown_until = 0.0 +_forecast_limit_daily = False # is the current forecast cooldown a daily-quota exhaustion? + class WeatherUnavailable(RuntimeError): """Weather data can't be fetched right now (rate limit, upstream outage with @@ -157,7 +193,8 @@ def _seconds_to_utc_reset() -> float: def _note_rate_limit(e) -> None: - """Record a rate limit: back off ~2 min, or until tomorrow for a daily quota.""" + """Record an archive rate limit: back off ~2 min, or until tomorrow for a + daily quota.""" global _archive_cooldown_until, _archive_limit_daily reason = _rate_limit_reason(e).lower() _archive_limit_daily = "daily" in reason or "tomorrow" in reason @@ -165,6 +202,16 @@ def _note_rate_limit(e) -> None: _seconds_to_utc_reset() if _archive_limit_daily else ARCHIVE_COOLDOWN) +def _note_forecast_rate_limit(e) -> None: + """Record a forecast-endpoint rate limit. Mirrors _note_rate_limit exactly, + but against the separate _forecast_cooldown_until (see its module comment).""" + global _forecast_cooldown_until, _forecast_limit_daily + reason = _rate_limit_reason(e).lower() + _forecast_limit_daily = "daily" in reason or "tomorrow" in reason + _forecast_cooldown_until = time.time() + ( + _seconds_to_utc_reset() if _forecast_limit_daily else FORECAST_COOLDOWN) + + def _cell_lock(cell_id: str) -> threading.Lock: with _LOCKS_GUARD: lk = _CELL_LOCKS.get(cell_id) @@ -176,11 +223,18 @@ def _cell_lock(cell_id: str) -> threading.Lock: def _request(url, params, timeout, *, phase, headers=None, attempts=MAX_ATTEMPTS): """GET with bounded retries; every retry and the final failure are logged to the errors folder (tagged ``retry`` / ``error``). A 429 (rate limit) fails fast - without retrying, so we don't hammer the limit and make it worse.""" + without retrying, so we don't hammer the limit and make it worse. + + ``timeout`` is either one number applied to every attempt (the common case), + or a sequence of per-attempt values (see ARCHIVE_FETCH_TIMEOUTS) so a caller + that holds a lock across every attempt can fail fast on the early ones and + reserve a longer wait for the last.""" last = None for attempt in range(1, attempts + 1): + attempt_timeout = (timeout[min(attempt - 1, len(timeout) - 1)] + if isinstance(timeout, (tuple, list)) else timeout) try: - r = httpx.get(url, params=params, timeout=timeout, headers=headers) + r = _client.get(url, params=params, timeout=attempt_timeout, headers=headers) r.raise_for_status() metrics.record_outbound(phase, "ok") return r @@ -283,9 +337,28 @@ def _with_doy(df: pl.DataFrame) -> pl.DataFrame: def _write_cache(df: pl.DataFrame, path: str) -> None: """Persist a daily record to the parquet cache — the raw record only, never - the derived doy column (it's recomputed at read time).""" - os.makedirs(CACHE_DIR, exist_ok=True) - df.drop("doy", strict=False).write_parquet(path, compression="zstd") + the derived doy column (it's recomputed at read time). + + Written to a tempfile in the same directory, then renamed into place + (os.replace is atomic on the same filesystem) — mirrors the pattern in + api/homepage.py's refresh() and data/places.py's _fetch. A reader (or a + second writer — warm_cities.py guards against overlapping runs, but a live + request racing a topup is normal) must never see a partially-written parquet + file; writing straight to ``path`` (the old behavior) could hand a truncated + file to a concurrent read.""" + directory = os.path.dirname(path) + os.makedirs(directory, exist_ok=True) + fd, tmp = tempfile.mkstemp(dir=directory, suffix=".tmp") + try: + os.close(fd) + df.drop("doy", strict=False).write_parquet(tmp, compression="zstd") + os.replace(tmp, path) + except BaseException: + try: + os.unlink(tmp) + except OSError: + pass + raise # --- cache backend dispatch -------------------------------------------------- @@ -462,7 +535,7 @@ def _to_frame(daily: dict) -> pl.DataFrame: def _fetch_history(cell: dict) -> pl.DataFrame: end = (datetime.date.today() - datetime.timedelta(days=ARCHIVE_LATENCY_DAYS)).isoformat() params = _om_daily_params(cell, start_date=START_DATE, end_date=end, models=ARCHIVE_MODEL) - r = _request(ARCHIVE_URL, params, 180, phase="history_fetch") + r = _request(ARCHIVE_URL, params, ARCHIVE_FETCH_TIMEOUTS, phase="history_fetch") return _to_frame(r.json()["daily"]) @@ -523,7 +596,7 @@ def _fetch_history_nasa(cell: dict) -> pl.DataFrame: "end": end, "format": "JSON", } - r = _request(NASA_POWER_URL, params, 180, phase="history_nasa") + r = _request(NASA_POWER_URL, params, ARCHIVE_FETCH_TIMEOUTS, phase="history_nasa") return _nasa_to_frame(r.json()["properties"]["parameter"]) @@ -798,14 +871,27 @@ def _load_recent_forecast(cell: dict) -> pl.DataFrame: "past_days": RECENT_PAST_DAYS, "forecast_days": FORECAST_DAYS, } - try: - r = _request(FORECAST_URL, params, 60, phase="recent_forecast_fetch") - df = _to_frame(r.json()["daily"]) - except Exception as e: # noqa: BLE001 - # Open-Meteo forecast unavailable (rate limit or outage). Fall back to MET - # Norway (yr.no) — global + keyless — mirroring the archive's NASA POWER - # backup. MET Norway is forecast-only (no recent past days), so this keeps - # the forecast / day-ahead views working in a degraded form. + # Skip Open-Meteo's forecast endpoint entirely while it's in its own + # rate-limit cooldown (see _note_forecast_rate_limit) -- mirrors + # _load_history's archive-cooldown check, so a forecast brownout doesn't + # retry-storm the endpoint from every subscribed cell and every live request + # independently; it falls straight through to the MET Norway backup instead. + df = None + primary_error = None + if time.time() >= _forecast_cooldown_until: + try: + r = _request(FORECAST_URL, params, 60, phase="recent_forecast_fetch") + df = _to_frame(r.json()["daily"]) + except Exception as e: # noqa: BLE001 + if is_rate_limit(e): + _note_forecast_rate_limit(e) + primary_error = e + + if df is None: + # Open-Meteo forecast unavailable (rate limit, cooldown, or outage). Fall + # back to MET Norway (yr.no) — global + keyless — mirroring the archive's + # NASA POWER backup. MET Norway is forecast-only (no recent past days), + # so this keeps the forecast / day-ahead views working in a degraded form. try: df = _fetch_forecast_metno(cell) except Exception: # noqa: BLE001 - backup unavailable too @@ -818,11 +904,16 @@ def _load_recent_forecast(cell: dict) -> pl.DataFrame: return _with_doy(hit[0]) # No cache either: surface the original error, classifying an Open-Meteo # rate limit as the typed, daily-aware WeatherUnavailable (the archive - # path classifies its own inside _load_history). - if is_rate_limit(e): - daily = "daily" in _rate_limit_reason(e).lower() - raise WeatherUnavailable(limit_message(daily), daily=daily) from e - raise + # path classifies its own inside _load_history). primary_error is None + # when the primary fetch was skipped outright (cooldown active) -- + # report the cooldown itself in that case. + if primary_error is not None and is_rate_limit(primary_error): + daily = "daily" in _rate_limit_reason(primary_error).lower() + raise WeatherUnavailable(limit_message(daily), daily=daily) from primary_error + if primary_error is not None: + raise primary_error + raise WeatherUnavailable(limit_message(_forecast_limit_daily), + daily=_forecast_limit_daily) _write_recent_backed(cell_id, df) return df @@ -830,13 +921,21 @@ def _load_recent_forecast(cell: dict) -> pl.DataFrame: _REVGEO_CACHE: dict[str, str | None] = {} # Nominatim's usage policy allows ~1 request/second and rejects bursts. The compare -# page loads several locations at once, so serialize the reverse lookups here and -# space them out — otherwise the burst is rate-limited, a null label gets cached, and -# those locations are left showing bare coordinates. -_REVGEO_LOCK = threading.Lock() +# page loads several locations at once, so the actual fetch + pacing run on ONE +# dedicated worker thread (see _revgeo_worker), never on a caller's own thread — +# in the server, a caller's thread is one of Starlette's shared sync threadpool +# threads, and the old design serialized AND slept (up to _REVGEO_MIN_INTERVAL) +# right there, pinning a threadpool thread for the whole wait. A single worker +# draining a queue gets the same ~1/sec pacing for free (nothing else ever calls +# Nominatim) without blocking anything but itself. _REVGEO_MIN_INTERVAL = 1.1 # seconds between successive Nominatim reverse calls _revgeo_last = 0.0 +_REVGEO_QUEUE: "queue.Queue[tuple[float, float, str, concurrent.futures.Future]]" = queue.Queue() +_REVGEO_WORKER_LOCK = threading.Lock() +_revgeo_worker_started = False +_REVGEO_WAIT_TIMEOUT = 10.0 # seconds a caller waits for ITS OWN request before giving up + def reverse_geocode_cached(lat: float, lon: float) -> tuple[bool, str | None]: """(found, label) from the in-memory + SQLite revgeo caches only — never calls @@ -850,64 +949,113 @@ def reverse_geocode_cached(lat: float, lon: float) -> tuple[bool, str | None]: return (found, label) +def _fetch_revgeo_label(lat: float, lon: float) -> str | None: + """The actual Nominatim reverse-geocode HTTP call + label assembly. Called + ONLY from _revgeo_worker (never on a caller's thread) — returns None (never + raises) so a bad response degrades to bare coordinates instead of killing the + worker thread.""" + try: + r = _request( + "https://nominatim.openstreetmap.org/reverse", + # zoom 14 resolves to the suburb/neighbourhood level so we can lead + # with it when OSM has one (zoom 10 only ever returns the city). + # accept-language=en keeps labels in one script worldwide (matches + # the forward geocoder's language=en). + {"lat": lat, "lon": lon, "format": "jsonv2", "zoom": 14, + "addressdetails": 1, "accept-language": "en"}, + 15, + phase="reverse_geocode", + headers={"User-Agent": "Thermograph/0.1 (local weather grading app)"}, + ) + a = r.json().get("address", {}) or {} + # Lead with the neighbourhood when available, then the city, then region. + neighborhood = (a.get("neighbourhood") or a.get("suburb") + or a.get("quarter") or a.get("city_district") + or a.get("borough")) + city = (a.get("city") or a.get("town") or a.get("village") + or a.get("hamlet") or a.get("county")) + region = a.get("state") or a.get("province") or a.get("region") + country = a.get("country") + # dict.fromkeys drops duplicates (e.g. neighbourhood == city) in order. + # Country trails the neighbourhood/city/region so the label reads as a full + # hierarchy ("West Seattle, Seattle, Washington, United States"); the + # frontend leads with the first part and mutes the rest. + parts = dict.fromkeys(p for p in (neighborhood, city, region, country) if p) + return ", ".join(parts) or None + except Exception: # noqa: BLE001 - reverse geocoding is a nicety, never fatal + return None + + +def _revgeo_worker() -> None: + """Drains _REVGEO_QUEUE one request at a time, forever. The sole caller of + _fetch_revgeo_label, so the ~1/sec pacing below is enforced just by doing the + work serially — no lock needed, since nothing else ever touches Nominatim. + Re-checks the cache before fetching (a request queued behind an identical one + is answered from what the earlier request just cached, no duplicate call), + and always finishes the fetch + persists the result even if the original + caller already gave up waiting (see reverse_geocode's timeout).""" + global _revgeo_last + while True: + lat, lon, key, fut = _REVGEO_QUEUE.get() + try: + found, label = reverse_geocode_cached(lat, lon) + if not found: + wait = _REVGEO_MIN_INTERVAL - (time.monotonic() - _revgeo_last) + if wait > 0: + time.sleep(wait) + label = _fetch_revgeo_label(lat, lon) + _revgeo_last = time.monotonic() + _REVGEO_CACHE[key] = label + store.put_revgeo(key, label) # survive restarts (a None label retries after its TTL) + if not fut.done(): + fut.set_result(label) + except Exception: # noqa: BLE001 - never let a bad request kill the worker + if not fut.done(): + fut.set_result(None) + finally: + _REVGEO_QUEUE.task_done() + + +def _start_revgeo_worker() -> None: + """Start the dedicated reverse-geocode worker thread, once per process + (idempotent, thread-safe). Lazy (on first use) rather than at import, since + climate.py has no app-lifespan hook of its own — mirrors data/places.py's + start_loading().""" + global _revgeo_worker_started + with _REVGEO_WORKER_LOCK: + if _revgeo_worker_started: + return + _revgeo_worker_started = True + threading.Thread(target=_revgeo_worker, name="revgeo-worker", daemon=True).start() + + def reverse_geocode(lat: float, lon: float) -> str | None: """Best-effort neighbourhood/city label for a point (OpenStreetMap Nominatim). Cached per ~cell — in memory for the process and in SQLite across restarts — so panning around (or redeploying) doesn't re-hit the service, and failures return None so the caller can fall back to bare coordinates. + + The actual fetch (and Nominatim's pacing) happens on the dedicated + _revgeo_worker thread, not here — this just enqueues the request and waits up + to _REVGEO_WAIT_TIMEOUT for its result, so a burst of concurrent uncached + lookups (the compare page fires several at once) queues behind the ~1/sec + limit without pinning a caller's (in the server, a shared threadpool) thread + for the wait. A timed-out wait returns None like any other failed lookup — + the caller falls back to bare coordinates — but the worker keeps going and + still caches the answer for the next call. """ key = store.revgeo_key(lat, lon) found, label = reverse_geocode_cached(lat, lon) if found: return label - # Serialize + rate-limit the upstream call. Concurrent callers (the compare page - # loads several locations at once) would otherwise burst past Nominatim's ~1/sec - # limit and get rate-limited, caching a null label. Under the lock we re-check the - # cache (a peer may have just resolved this cell) and space successive calls out. - global _revgeo_last - with _REVGEO_LOCK: - found, label = reverse_geocode_cached(lat, lon) - if found: - return label - wait = _REVGEO_MIN_INTERVAL - (time.monotonic() - _revgeo_last) - if wait > 0: - time.sleep(wait) - label = None - try: - r = _request( - "https://nominatim.openstreetmap.org/reverse", - # zoom 14 resolves to the suburb/neighbourhood level so we can lead - # with it when OSM has one (zoom 10 only ever returns the city). - # accept-language=en keeps labels in one script worldwide (matches - # the forward geocoder's language=en). - {"lat": lat, "lon": lon, "format": "jsonv2", "zoom": 14, - "addressdetails": 1, "accept-language": "en"}, - 15, - phase="reverse_geocode", - headers={"User-Agent": "Thermograph/0.1 (local weather grading app)"}, - ) - a = r.json().get("address", {}) or {} - # Lead with the neighbourhood when available, then the city, then region. - neighborhood = (a.get("neighbourhood") or a.get("suburb") - or a.get("quarter") or a.get("city_district") - or a.get("borough")) - city = (a.get("city") or a.get("town") or a.get("village") - or a.get("hamlet") or a.get("county")) - region = a.get("state") or a.get("province") or a.get("region") - country = a.get("country") - # dict.fromkeys drops duplicates (e.g. neighbourhood == city) in order. - # Country trails the neighbourhood/city/region so the label reads as a full - # hierarchy ("West Seattle, Seattle, Washington, United States"); the - # frontend leads with the first part and mutes the rest. - parts = dict.fromkeys(p for p in (neighborhood, city, region, country) if p) - label = ", ".join(parts) or None - except Exception: # noqa: BLE001 - reverse geocoding is a nicety, never fatal - label = None - _revgeo_last = time.monotonic() - _REVGEO_CACHE[key] = label - store.put_revgeo(key, label) # survive restarts (a None label retries after its TTL) - return label + _start_revgeo_worker() + fut: "concurrent.futures.Future[str | None]" = concurrent.futures.Future() + _REVGEO_QUEUE.put((lat, lon, key, fut)) + try: + return fut.result(timeout=_REVGEO_WAIT_TIMEOUT) + except concurrent.futures.TimeoutError: + return None def geocode(name: str, count: int = 5) -> list[dict]: diff --git a/data/climate_store.py b/data/climate_store.py index 69d77fb..19c883c 100644 --- a/data/climate_store.py +++ b/data/climate_store.py @@ -24,6 +24,7 @@ recompute-from-parquet), a miss here degrades to *fetch-from-upstream* — prod no parquet fallback, and ``climate.py``'s loaders already treat "no cache" as "go fetch". """ +import contextlib import os import threading @@ -68,29 +69,69 @@ def _libpq_dsn(url: str) -> str: _PG_DSN = _libpq_dsn(DATABASE_URL) if IS_POSTGRES else "" -# Per-thread connection: uvicorn serves on a thread pool and psycopg connections -# aren't thread-safe, so each thread lazily opens its own (mirrors store.py). -_local = threading.local() +# Bounded pool, replacing the old thread-local connection (one psycopg.connect per +# THREAD, opened lazily, never closed). Starlette's sync threadpool can spin up +# dozens of worker threads per uvicorn worker; with N workers x this module x +# data/store.py all doing the same thing, that had no ceiling and could exhaust +# Postgres max_connections, taking down every module sharing the database at +# once. min_size=1 keeps an idle deploy cheap; max_size=8 is a modest per-process +# ceiling that still lets several concurrent cell fetches proceed without +# serializing on one connection. connect_timeout + statement_timeout bound how +# long a caller (or, worse, the notifier's leader-election flock) can be wedged +# by an unresponsive Postgres. +_POOL_MIN_SIZE = 1 +_POOL_MAX_SIZE = 8 + +_pool_lock = threading.Lock() +_pool_obj = None # type: "psycopg_pool.ConnectionPool | None" -def _conn(): - """Thread-local autocommit psycopg connection, or None when Postgres is off or - unreachable. A dropped connection is retired so the next call reconnects.""" +def _pool(): + """Lazily open (once) the bounded connection pool, or None when Postgres is + off or the pool can't be opened — callers treat that exactly like the old + _conn() returning None.""" + global _pool_obj if not IS_POSTGRES: return None - conn = getattr(_local, "conn", None) - if conn is not None: - if getattr(conn, "closed", False): - _local.conn = None - else: - return conn - try: - import psycopg # local import: only the Postgres path needs it - conn = psycopg.connect(_PG_DSN, autocommit=True) - _local.conn = conn - return conn - except Exception: # noqa: BLE001 - the store is optional; callers fetch on miss - return None + if _pool_obj is not None: + return _pool_obj + with _pool_lock: + if _pool_obj is None: + try: + import psycopg_pool # local import: only the Postgres path needs it + _pool_obj = psycopg_pool.ConnectionPool( + _PG_DSN, + min_size=_POOL_MIN_SIZE, + max_size=_POOL_MAX_SIZE, + kwargs={ + "autocommit": True, + "connect_timeout": 5, + "options": "-c statement_timeout=30000", + }, + open=True, + ) + except Exception: # noqa: BLE001 - the store is optional; callers fetch on miss + return None + return _pool_obj + + +@contextlib.contextmanager +def _conn(): + """Borrow a connection from the bounded pool for the duration of one + operation, or yield None when Postgres is off or unreachable. Every caller + frames exactly one operation in a ``with _conn() as conn:`` block, so the + connection is returned to the pool (or discarded, if broken) the moment that + operation finishes — replacing the old pattern where a thread opened one raw + connection and held it open, unclosed, for its entire lifetime.""" + if not IS_POSTGRES: + yield None + return + pool = _pool() + if pool is None: + yield None + return + with pool.connection() as conn: + yield conn # --- freshness (climate_sync) ----------------------------------------------- @@ -124,12 +165,12 @@ def _bump_sync(conn, cell_id, *, history=None, recent=None, cols_version=None) - def history_synced_at(cell_id: str) -> float: """Epoch seconds of the cell's last history write/topup; 0.0 if absent.""" - conn = _conn() - if conn is None: - return 0.0 try: - row = _get_sync(conn, cell_id) - return float(row[0]) if row and row[0] is not None else 0.0 + with _conn() as conn: + if conn is None: + return 0.0 + row = _get_sync(conn, cell_id) + return float(row[0]) if row and row[0] is not None else 0.0 except Exception: # noqa: BLE001 return 0.0 @@ -137,24 +178,24 @@ def history_synced_at(cell_id: str) -> float: def recent_synced_at(cell_id: str) -> float: """Epoch seconds of the cell's last recent-bundle write; 0.0 if absent. Backs climate.recent_stamp, so it must advance only on a real rewrite (see write_recent).""" - conn = _conn() - if conn is None: - return 0.0 try: - row = _get_sync(conn, cell_id) - return float(row[1]) if row and row[1] is not None else 0.0 + with _conn() as conn: + if conn is None: + return 0.0 + row = _get_sync(conn, cell_id) + return float(row[1]) if row and row[1] is not None else 0.0 except Exception: # noqa: BLE001 return 0.0 def cols_version(cell_id: str) -> int: """The stored schema version for a cell's history; 0 if absent.""" - conn = _conn() - if conn is None: - return 0 try: - row = _get_sync(conn, cell_id) - return int(row[2]) if row and row[2] is not None else 0 + with _conn() as conn: + if conn is None: + return 0 + row = _get_sync(conn, cell_id) + return int(row[2]) if row and row[2] is not None else 0 except Exception: # noqa: BLE001 return 0 @@ -175,19 +216,19 @@ def read_history(cell_id: str) -> "tuple[pl.DataFrame, float] | None": """(history frame, history_synced_at) for a cell, or None when there is no schema-complete cached record. A stored cols_version below COLS_VERSION reads as absent, so climate.py refetches to add the new column(s).""" - conn = _conn() - if conn is None: - return None try: - sync = _get_sync(conn, cell_id) - if sync is None or sync[0] is None: - return None - if (sync[2] or 0) < COLS_VERSION: - return None - df = _read_frame(conn, "climate_history", cell_id) - if df.is_empty(): - return None - return df, float(sync[0]) + with _conn() as conn: + if conn is None: + return None + sync = _get_sync(conn, cell_id) + if sync is None or sync[0] is None: + return None + if (sync[2] or 0) < COLS_VERSION: + return None + df = _read_frame(conn, "climate_history", cell_id) + if df.is_empty(): + return None + return df, float(sync[0]) except Exception: # noqa: BLE001 return None @@ -197,29 +238,29 @@ def read_history_raw(cell_id: str) -> "pl.DataFrame | None": fallback when every upstream fetch fails (the Postgres analogue of reading the parquet file directly, bypassing the NEW_COLS completeness check). None when there are no rows.""" - conn = _conn() - if conn is None: - return None try: - df = _read_frame(conn, "climate_history", cell_id) - return None if df.is_empty() else df + with _conn() as conn: + if conn is None: + return None + df = _read_frame(conn, "climate_history", cell_id) + return None if df.is_empty() else df except Exception: # noqa: BLE001 return None def read_recent(cell_id: str) -> "tuple[pl.DataFrame, float] | None": """(recent+forecast frame, recent_synced_at) for a cell, or None when absent.""" - conn = _conn() - if conn is None: - return None try: - sync = _get_sync(conn, cell_id) - if sync is None or sync[1] is None: - return None - df = _read_frame(conn, "climate_recent", cell_id) - if df.is_empty(): - return None - return df, float(sync[1]) + with _conn() as conn: + if conn is None: + return None + sync = _get_sync(conn, cell_id) + if sync is None or sync[1] is None: + return None + df = _read_frame(conn, "climate_recent", cell_id) + if df.is_empty(): + return None + return df, float(sync[1]) except Exception: # noqa: BLE001 return None @@ -238,21 +279,21 @@ def write_history(cell_id: str, df: pl.DataFrame, synced_at: float) -> None: """Upsert a cell's history rows and stamp history_synced_at. Idempotent: a re-fetched day supersedes its stored duplicate (ON CONFLICT DO UPDATE), which mirrors the parquet unique(subset="date", keep="last") merge. Best-effort.""" - conn = _conn() - if conn is None: - return try: - with conn.transaction(): - conn.execute( - "CREATE TEMP TABLE _hist_stage (LIKE climate_history) ON COMMIT DROP") - with conn.cursor() as cur: - _copy_rows(cur, "_hist_stage", cell_id, df) - set_cols = ", ".join(f"{c} = EXCLUDED.{c}" for c in COLS if c != "date") - conn.execute( - f"INSERT INTO climate_history (cell_id, {', '.join(COLS)}) " - f"SELECT cell_id, {', '.join(COLS)} FROM _hist_stage " - f"ON CONFLICT (cell_id, date) DO UPDATE SET {set_cols}") - _bump_sync(conn, cell_id, history=synced_at, cols_version=COLS_VERSION) + with _conn() as conn: + if conn is None: + return + with conn.transaction(): + conn.execute( + "CREATE TEMP TABLE _hist_stage (LIKE climate_history) ON COMMIT DROP") + with conn.cursor() as cur: + _copy_rows(cur, "_hist_stage", cell_id, df) + set_cols = ", ".join(f"{c} = EXCLUDED.{c}" for c in COLS if c != "date") + conn.execute( + f"INSERT INTO climate_history (cell_id, {', '.join(COLS)}) " + f"SELECT cell_id, {', '.join(COLS)} FROM _hist_stage " + f"ON CONFLICT (cell_id, date) DO UPDATE SET {set_cols}") + _bump_sync(conn, cell_id, history=synced_at, cols_version=COLS_VERSION) except Exception: # noqa: BLE001 - failing to cache must not fail the request pass @@ -261,11 +302,11 @@ def touch_history(cell_id: str, synced_at: float) -> None: """Bump only history_synced_at — the analogue of ``os.utime`` on the parquet file when the archive tail is already current, so the hourly topup timer resets without a rewrite. Best-effort.""" - conn = _conn() - if conn is None: - return try: - _bump_sync(conn, cell_id, history=synced_at) + with _conn() as conn: + if conn is None: + return + _bump_sync(conn, cell_id, history=synced_at) except Exception: # noqa: BLE001 pass @@ -276,14 +317,14 @@ def write_recent(cell_id: str, df: pl.DataFrame, synced_at: float) -> None: refresh, so a full replace is the simplest correct write. Call ONLY on a real rewrite — recent_synced_at feeds recent_stamp, which must not advance when a stale bundle is served without a refetch. Best-effort.""" - conn = _conn() - if conn is None: - return try: - with conn.transaction(): - conn.execute("DELETE FROM climate_recent WHERE cell_id = %s", (cell_id,)) - with conn.cursor() as cur: - _copy_rows(cur, "climate_recent", cell_id, df) - _bump_sync(conn, cell_id, recent=synced_at) + with _conn() as conn: + if conn is None: + return + with conn.transaction(): + conn.execute("DELETE FROM climate_recent WHERE cell_id = %s", (cell_id,)) + with conn.cursor() as cur: + _copy_rows(cur, "climate_recent", cell_id, df) + _bump_sync(conn, cell_id, recent=synced_at) except Exception: # noqa: BLE001 pass diff --git a/data/places.py b/data/places.py index c2faadb..c80a9f1 100644 --- a/data/places.py +++ b/data/places.py @@ -34,6 +34,13 @@ import paths GEO_DIR = os.path.join(paths.DATA_DIR, "geonames") GEONAMES_URL = "https://download.geonames.org/export/dump/" +# Reused connection instead of a bare httpx.get per call -- mirrors +# web/app.py's reused _frontend_client / data/climate.py's _client. Only a +# handful of calls ever happen (the dump files are cached forever, see _fetch), +# but a shared client costs nothing and keeps the pattern consistent everywhere +# this module talks HTTP. +_client = httpx.Client() + # Which GeoNames cities dump to index. cities1000 (~170k places, population # ≥ 1000) is small enough to hold in memory and big enough to cover the tiny # vacation towns people actually search for; set THERMOGRAPH_CITIES=cities5000 @@ -74,7 +81,7 @@ def _fetch(name: str) -> str: path = os.path.join(GEO_DIR, name) if os.path.exists(path) and os.path.getsize(path) > 0: return path - r = httpx.get(GEONAMES_URL + name, timeout=120, follow_redirects=True) + r = _client.get(GEONAMES_URL + name, timeout=120, follow_redirects=True) r.raise_for_status() tmp = path + ".part" with open(tmp, "wb") as f: diff --git a/data/store.py b/data/store.py index 4ebc288..d025994 100644 --- a/data/store.py +++ b/data/store.py @@ -21,9 +21,11 @@ their own errors), and deleting data/thermograph.sqlite is a safe full reset. On Postgres (``THERMOGRAPH_DATABASE_URL=postgresql://…``) the same two tables live as UNLOGGED tables — fast and non-durable, which suits a rebuildable cache — and -the per-thread connection uses psycopg instead of sqlite3. Every fail-soft path is -identical: a connection or query that fails just degrades to recompute-from-parquet. +psycopg (via a bounded pool, see _get_pg_pool) replaces sqlite3. Every fail-soft +path is identical: a connection or query that fails just degrades to +recompute-from-parquet. """ +import contextlib import json import os import sqlite3 @@ -118,61 +120,122 @@ else: _SQL_PUT_REVGEO = "INSERT OR REPLACE INTO revgeo VALUES (?,?,?)" # SQLite connections aren't shareable across threads; uvicorn serves requests on -# a thread pool, so each thread lazily opens its own connection. WAL lets those -# readers run concurrently with the (serialized) writers. The Postgres path keeps -# the same per-thread model (psycopg connections likewise aren't thread-safe). +# a thread pool, so each thread lazily opens (and keeps, for its lifetime) its own +# connection. WAL lets those readers run concurrently with the (serialized) +# writers. There's no server-side connection limit to protect against here, so +# this thread-local model stays as-is for SQLite. _local = threading.local() +# Postgres instead uses a bounded pool (see _get_pg_pool), not a thread-local +# connection: the old pattern opened one raw psycopg connection per THREAD, +# lazily, and never closed it. Starlette's sync threadpool can spin up dozens of +# worker threads per uvicorn worker; with N workers x this module x +# data/climate_store.py doing the same thing, that had no ceiling and could +# exhaust Postgres max_connections. min_size=1 keeps an idle deploy cheap; +# max_size=8 is a modest per-process ceiling that still lets several concurrent +# requests proceed without serializing on one connection. connect_timeout + +# statement_timeout bound how long a caller (or the notifier's leader-election +# flock) can be wedged by an unresponsive Postgres. +_pg_pool_lock = threading.Lock() +_pg_pool = None # type: "psycopg_pool.ConnectionPool | None" -def _open_pg(): - """Open a thread-local psycopg connection and ensure the schema. autocommit so - reads never leave an idle-in-transaction and each write lands immediately (the - shared put/get code still calls ``.commit()``, a harmless no-op under autocommit).""" + +def _configure_pg(conn) -> None: + """Run once per NEW physical connection the pool opens (psycopg_pool's + ``configure=``), not on every checkout — mirrors the old _open_pg's + connect-then-ensure-schema sequence without re-running the (idempotent) + DDL on every borrow. + + Multiple connections can legitimately open around the same moment (the + pool's min_size warm-up racing an immediate caller's on-demand open; two + uvicorn workers' pools both cold-starting at once) and each runs this same + CREATE TABLE IF NOT EXISTS. Despite IF NOT EXISTS, two sessions creating the + same brand-new table at the same instant can still hit a duplicate-key error + on Postgres's catalog — a documented race, not a real failure: it just means + the table exists either way, so it's swallowed here rather than failing this + connection's open (which put_payload/get_payload etc. would otherwise + silently treat as "store unavailable" and drop the operation).""" import psycopg # local import: only the Postgres path needs it - conn = psycopg.connect(_PG_DSN, autocommit=True) for stmt in _PG_SCHEMA: - conn.execute(stmt) - return conn + try: + conn.execute(stmt) + except psycopg.errors.UniqueViolation: + pass +def _get_pg_pool(): + """Lazily open (once) the bounded Postgres pool, or None if it can't be + opened — callers treat that exactly like the old _conn() returning None.""" + global _pg_pool + if _pg_pool is not None: + return _pg_pool + with _pg_pool_lock: + if _pg_pool is None: + try: + import psycopg_pool # local import: only the Postgres path needs it + _pg_pool = psycopg_pool.ConnectionPool( + _PG_DSN, + min_size=1, + max_size=8, + kwargs={ + "autocommit": True, + "connect_timeout": 5, + "options": "-c statement_timeout=30000", + }, + configure=_configure_pg, + open=True, + ) + except Exception: # noqa: BLE001 - the store is optional; readers fall back + return None + return _pg_pool + + +@contextlib.contextmanager def _conn(): + """Yield a working connection, or None when the store is unavailable. + + Postgres: borrowed from the bounded pool above for the duration of one + operation and returned automatically on exit. SQLite: the thread's + persistent connection (opened once, kept for the thread's lifetime, same as + before). Either way every caller frames exactly one operation in a + ``with _conn() as conn:`` block.""" + if IS_POSTGRES: + pool = _get_pg_pool() + if pool is None: + yield None + return + with pool.connection() as conn: + yield conn + return conn = getattr(_local, "conn", None) - if conn is not None: - # A dropped Postgres connection would otherwise disable the cache for the - # life of the thread; retire it so the next call reconnects. - if IS_POSTGRES and getattr(conn, "closed", False): - _local.conn = None - else: - return conn - try: - if IS_POSTGRES: - conn = _open_pg() - else: + if conn is None: + try: os.makedirs(os.path.dirname(DB_PATH), exist_ok=True) conn = sqlite3.connect(DB_PATH, timeout=5.0) conn.execute("PRAGMA journal_mode=WAL") conn.execute("PRAGMA synchronous=NORMAL") conn.executescript(_SCHEMA) - _local.conn = conn - return conn - except Exception: # noqa: BLE001 - the store is optional; readers fall back - return None + _local.conn = conn + except Exception: # noqa: BLE001 - the store is optional; readers fall back + yield None + return + yield conn # --- derived payloads ------------------------------------------------------- def get_payload(kind: str, cell_id: str, key: str, token: str) -> bytes | None: """Raw JSON bytes of a cached payload, or None when absent or token-invalid.""" - conn = _conn() - if conn is None: - return None try: - row = conn.execute(_SQL_GET_PAYLOAD, (kind, cell_id, key)).fetchone() - if row is None or row[0] != token: - return None - # psycopg returns bytea as bytes/memoryview; bytes() normalizes both. - data = bytes(row[1]) if IS_POSTGRES else row[1] - return zlib.decompress(data) + with _conn() as conn: + if conn is None: + return None + row = conn.execute(_SQL_GET_PAYLOAD, (kind, cell_id, key)).fetchone() + if row is None or row[0] != token: + return None + # psycopg returns bytea as bytes/memoryview; bytes() normalizes both. + data = bytes(row[1]) if IS_POSTGRES else row[1] + return zlib.decompress(data) except Exception: # noqa: BLE001 return None @@ -196,14 +259,15 @@ def put_payload(kind: str, cell_id: str, key: str, token: str, payload: dict, grading should fail loudly here, not be cached as JS-unparseable `NaN`.""" body = json.dumps(payload, default=str, allow_nan=False, separators=(",", ":")).encode() - conn = _conn() - if cache and conn is not None: + if cache: try: - conn.execute( - _SQL_PUT_PAYLOAD, - (kind, cell_id, key, token, zlib.compress(body, 6), time.time()), - ) - conn.commit() + with _conn() as conn: + if conn is not None: + conn.execute( + _SQL_PUT_PAYLOAD, + (kind, cell_id, key, token, zlib.compress(body, 6), time.time()), + ) + conn.commit() except Exception: # noqa: BLE001 - failing to cache must not fail the request pass return body @@ -223,43 +287,43 @@ def revgeo_key(lat: float, lon: float) -> str: def get_revgeo(key: str) -> tuple[bool, str | None]: """(found, label). A stored NULL label counts as found (a cached miss) until it ages past REVGEO_MISS_TTL, when the lookup becomes retryable.""" - conn = _conn() - if conn is None: - return (False, None) try: - row = conn.execute(_SQL_GET_REVGEO, (key,)).fetchone() - if row is None: - return (False, None) - if row[0] is None and time.time() - row[1] > REVGEO_MISS_TTL: - return (False, None) - return (True, row[0]) + with _conn() as conn: + if conn is None: + return (False, None) + row = conn.execute(_SQL_GET_REVGEO, (key,)).fetchone() + if row is None: + return (False, None) + if row[0] is None and time.time() - row[1] > REVGEO_MISS_TTL: + return (False, None) + return (True, row[0]) except Exception: # noqa: BLE001 return (False, None) def put_revgeo(key: str, label: str | None) -> None: - conn = _conn() - if conn is None: - return try: - conn.execute(_SQL_PUT_REVGEO, (key, label, time.time())) - conn.commit() + with _conn() as conn: + if conn is None: + return + conn.execute(_SQL_PUT_REVGEO, (key, label, time.time())) + conn.commit() except Exception: # noqa: BLE001 pass def stats() -> dict: """Row counts + on-disk size, for the migrate script's summary output.""" - conn = _conn() out = {"db_path": os.path.abspath(DB_PATH), "derived": 0, "revgeo": 0, "bytes": 0} - if conn is None: - return out try: - out["derived"] = conn.execute("SELECT COUNT(*) FROM derived").fetchone()[0] - out["revgeo"] = conn.execute("SELECT COUNT(*) FROM revgeo").fetchone()[0] - # Recent writes may still live in the WAL sidecar — count both files. - out["bytes"] = sum(os.path.getsize(p) for p in (DB_PATH, DB_PATH + "-wal") - if os.path.exists(p)) + with _conn() as conn: + if conn is None: + return out + out["derived"] = conn.execute("SELECT COUNT(*) FROM derived").fetchone()[0] + out["revgeo"] = conn.execute("SELECT COUNT(*) FROM revgeo").fetchone()[0] + # Recent writes may still live in the WAL sidecar — count both files. + out["bytes"] = sum(os.path.getsize(p) for p in (DB_PATH, DB_PATH + "-wal") + if os.path.exists(p)) except Exception: # noqa: BLE001 pass return out diff --git a/requirements.txt b/requirements.txt index b8bb685..c12d96d 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,11 +6,13 @@ numpy==2.2.1 # Accounts + notification subscriptions (see accounts/db.py, users.py, notify.py). # Postgres in prod/containers: asyncpg (async web/store) + psycopg (sync notifier); # aiosqlite is kept for the SQLite fallback the test suite runs on. alembic manages -# the accounts schema. +# the accounts schema. The `pool` extra pulls in psycopg_pool, used by +# data/climate_store.py + data/store.py to bound the raw connections those modules +# open on Postgres (see their module docstrings) instead of one-per-thread forever. fastapi-users[sqlalchemy]==15.0.5 aiosqlite==0.22.1 asyncpg==0.30.0 -psycopg[binary]==3.2.3 +psycopg[binary,pool]==3.2.3 alembic==1.14.0 # Web Push (VAPID) delivery of notifications (see notifications/push.py). Pulls in # py-vapid, cryptography, and http-ece. diff --git a/tests/accounts/test_db_engines.py b/tests/accounts/test_db_engines.py index 8bc5224..29454d5 100644 --- a/tests/accounts/test_db_engines.py +++ b/tests/accounts/test_db_engines.py @@ -28,6 +28,31 @@ def test_sync_url_derivation(): assert db._sync_url("sqlite+aiosqlite:////tmp/x.sqlite") == "sqlite:////tmp/x.sqlite" +def test_postgres_pool_sizing_and_sync_timeouts(): + """The Postgres engine-construction kwargs, defined at module scope so they're + testable without a real Postgres connection (the engines built from them are + only actually constructed when IS_POSTGRES, exercised by the docker-compose + stack -- see the module docstring). pool_size=1/max_overflow=1 (the old + setting) let a 3rd concurrent request per worker wait out pool_timeout for a + slot; both async engines need real headroom, and the sync engine (the + notifier's, driven off-event-loop) must bound connect + statement time so a + wedged Postgres can't hang it forever under its leader flock.""" + assert db._ASYNC_POOL_KWARGS["pool_size"] >= 5 + assert db._ASYNC_POOL_KWARGS["max_overflow"] >= 5 + + assert db._SYNC_POOL_KWARGS["pool_size"] < db._ASYNC_POOL_KWARGS["pool_size"] + connect_args = db._SYNC_POOL_KWARGS["connect_args"] + assert connect_args["connect_timeout"] > 0 + assert "statement_timeout" in connect_args["options"] + + # Build a real (unconnected -- SQLAlchemy engines are lazy) engine from these + # exact kwargs to confirm they're accepted by the psycopg dialect. + import sqlalchemy + engine = sqlalchemy.create_engine("postgresql+psycopg://u:p@h:5432/d", **db._SYNC_POOL_KWARGS) + assert engine.pool.size() == db._SYNC_POOL_KWARGS["pool_size"] + engine.dispose() + + def test_read_session_yields_a_working_session(): # On SQLite the read session is fully usable (no read-only enforcement); this # just confirms the dependency wiring resolves to a live AsyncSession. diff --git a/tests/data/test_climate.py b/tests/data/test_climate.py index e8e5782..fe4dd4d 100644 --- a/tests/data/test_climate.py +++ b/tests/data/test_climate.py @@ -324,3 +324,200 @@ def test_full_archive_is_accepted(monkeypatch, tmp_path): df, meta = climate._load_history(cell) assert meta["source"] == "open-meteo" assert df.height == full.height + + +# --- forecast cooldown (mirrors the archive path's, tracked separately) ----- + +def test_forecast_cooldown_skips_open_meteo_and_falls_to_metno(monkeypatch, tmp_path): + """While the forecast cooldown is active, the forecast fetch must not call + Open-Meteo at all -- straight to the MET Norway backup, mirroring + _load_history's archive-cooldown skip.""" + import time as time_mod + monkeypatch.setattr(climate, "CACHE_DIR", str(tmp_path)) + monkeypatch.setattr(climate, "_forecast_cooldown_until", time_mod.time() + 60) + cell = {"id": "fc_cooldown_cell", "center_lat": 47.6, "center_lon": -122.3} + + met = {"timeseries": [_metno_step("2026-07-16T00:00:00Z", 18.0, 55.0, 3.0, p1=0.0)]} + + class Resp: + def json(self): return {"properties": met} + + def fake_request(url, params, timeout, *, phase, headers=None, attempts=climate.MAX_ATTEMPTS): + assert url != climate.FORECAST_URL, "must not call Open-Meteo during cooldown" + assert url == climate.METNO_URL + return Resp() + + monkeypatch.setattr(climate, "_request", fake_request) + df = climate._load_recent_forecast(cell) + assert len(df) == 1 + + +def test_forecast_429_sets_its_own_cooldown(monkeypatch, tmp_path): + """A 429 from the forecast endpoint sets _forecast_cooldown_until (NOT + _archive_cooldown_until -- the two upstream endpoints have independent + quotas) and surfaces as WeatherUnavailable when the backup is also down.""" + import time as time_mod + monkeypatch.setattr(climate, "CACHE_DIR", str(tmp_path)) + monkeypatch.setattr(climate, "_forecast_cooldown_until", 0.0) + monkeypatch.setattr(climate, "_archive_cooldown_until", 0.0) + cell = {"id": "fc_429_cell", "center_lat": 47.6, "center_lon": -122.3} + + class FakeResponse: + status_code = 429 + def json(self): return {"reason": "rate limited"} + + class FakeRateLimitError(Exception): + def __init__(self): + self.response = FakeResponse() + + def fake_request(url, params, timeout, *, phase, headers=None, attempts=climate.MAX_ATTEMPTS): + if url == climate.FORECAST_URL: + raise FakeRateLimitError() + raise RuntimeError("metno also down") + + monkeypatch.setattr(climate, "_request", fake_request) + with pytest.raises(climate.WeatherUnavailable): + climate._load_recent_forecast(cell) + assert climate._forecast_cooldown_until > time_mod.time() + assert climate._archive_cooldown_until == 0.0 # the archive cooldown is untouched + + +def test_forecast_cooldown_does_not_block_archive_fetch(monkeypatch, tmp_path): + """The two cooldowns are independent state: an active forecast cooldown must + not gate _load_history's archive fetch.""" + import time as time_mod + monkeypatch.setattr(climate, "CACHE_DIR", str(tmp_path)) + monkeypatch.setattr(climate, "_forecast_cooldown_until", time_mod.time() + 60) + monkeypatch.setattr(climate, "_archive_cooldown_until", 0.0) + cell = {"id": "fc_indep_cell", "center_lat": 47.6, "center_lon": -122.3} + + full = _full_history_frame() + monkeypatch.setattr(climate, "_fetch_history", lambda c: full) + def _nasa_should_not_run(c): raise AssertionError("NASA should not be called") + monkeypatch.setattr(climate, "_fetch_history_nasa", _nasa_should_not_run) + + df, meta = climate._load_history(cell) + assert meta["source"] == "open-meteo" + + +# --- per-attempt timeouts (the cell-lock finding) --------------------------- + +def test_request_applies_a_per_attempt_timeout_sequence(monkeypatch): + """A timeout sequence (e.g. ARCHIVE_FETCH_TIMEOUTS) is applied per attempt, + not the same value on every retry -- the early attempts fail fast and only + the last gets the longer allowance, bounding how long a caller holding a + lock across every attempt (see _load_history) can be pinned.""" + import time as time_mod + seen = [] + + class FakeResp: + def raise_for_status(self): + pass + + class FakeClient: + def get(self, url, params=None, timeout=None, headers=None): + seen.append(timeout) + if len(seen) < 3: + raise RuntimeError("simulated transient failure") + return FakeResp() + + monkeypatch.setattr(climate, "_client", FakeClient()) + monkeypatch.setattr(time_mod, "sleep", lambda s: None) # skip retry backoff + climate._request("http://example.invalid", {}, (60, 60, 150), phase="test") + assert seen == [60, 60, 150] + + +def test_request_scalar_timeout_is_unchanged_for_every_attempt(monkeypatch): + """Every other _request caller passes a single number -- confirm that path is + untouched (same value on every attempt, not just the first).""" + import time as time_mod + seen = [] + + class FakeResp: + def raise_for_status(self): + pass + + class FakeClient: + def get(self, url, params=None, timeout=None, headers=None): + seen.append(timeout) + if len(seen) < 3: + raise RuntimeError("simulated transient failure") + return FakeResp() + + monkeypatch.setattr(climate, "_client", FakeClient()) + monkeypatch.setattr(time_mod, "sleep", lambda s: None) + climate._request("http://example.invalid", {}, 30, phase="test") + assert seen == [30, 30, 30] + + +# --- atomic parquet writes (warm_cities.py overlap-run finding) ------------- + +def test_write_cache_leaves_no_tempfile_behind_on_success(tmp_path): + df = climate._to_frame(_om_daily()) + path = str(tmp_path / "cell.parquet") + climate._write_cache(df, path) + import os + assert sorted(os.listdir(tmp_path)) == ["cell.parquet"] + + +def test_write_cache_cleans_up_tempfile_and_leaves_no_partial_target_on_failure( + monkeypatch, tmp_path): + """A write failure (disk full, killed process, ...) must not leave a + half-written file at the real path, and must not leak the tempfile either -- + a concurrent reader (or an overlapping warm_cities.py run) must only ever see + the old complete file or the new complete file, never a partial one.""" + df = climate._to_frame(_om_daily()) + path = str(tmp_path / "cell.parquet") + + def boom(self, *a, **kw): + raise RuntimeError("disk full") + + monkeypatch.setattr(pl.DataFrame, "write_parquet", boom) + with pytest.raises(RuntimeError): + climate._write_cache(df, path) + + import os + assert not os.path.exists(path) + assert os.listdir(tmp_path) == [] + + +# --- reverse-geocode off the shared threadpool ------------------------------ + +def test_reverse_geocode_cache_hit_never_touches_the_worker_queue(monkeypatch): + key = climate.store.revgeo_key(10.0, 20.0) + monkeypatch.setitem(climate._REVGEO_CACHE, key, "Cached Place") + + def boom(*a, **kw): + raise AssertionError("a cache hit must not enqueue a worker request") + + monkeypatch.setattr(climate._REVGEO_QUEUE, "put", boom) + assert climate.reverse_geocode(10.0, 20.0) == "Cached Place" + + +def test_reverse_geocode_miss_resolves_via_the_worker_thread(monkeypatch): + """An uncached lookup is answered by the dedicated worker thread (not the + calling thread), and the result is cached for the next call.""" + monkeypatch.setattr(climate, "_fetch_revgeo_label", lambda lat, lon: "Worker Place") + label = climate.reverse_geocode(11.111, 22.222) + assert label == "Worker Place" + found, cached = climate.reverse_geocode_cached(11.111, 22.222) + assert found and cached == "Worker Place" + + +def test_reverse_geocode_timeout_returns_none_without_blocking_the_caller(monkeypatch): + """A caller waits only up to _REVGEO_WAIT_TIMEOUT for ITS OWN request, even if + the worker is still busy on it -- it must not block for as long as the fetch + itself takes (that was exactly the old threadpool-pinning problem).""" + import time as time_mod + monkeypatch.setattr(climate, "_REVGEO_WAIT_TIMEOUT", 0.05) + + def slow_fetch(lat, lon): + time_mod.sleep(0.3) + return "Too Slow" + + monkeypatch.setattr(climate, "_fetch_revgeo_label", slow_fetch) + t0 = time_mod.monotonic() + label = climate.reverse_geocode(33.333, 44.444) + elapsed = time_mod.monotonic() - t0 + assert label is None + assert elapsed < 0.2 # returned near the wait timeout, not after the 0.3s fetch diff --git a/tests/data/test_climate_store.py b/tests/data/test_climate_store.py index b31bf35..a6fbcd3 100644 --- a/tests/data/test_climate_store.py +++ b/tests/data/test_climate_store.py @@ -81,11 +81,10 @@ def pg_store(): pytest.skip("set THERMOGRAPH_TEST_DATABASE_URL to run the Postgres bridge test") import psycopg - saved = (climate_store.IS_POSTGRES, climate_store._PG_DSN) + saved = (climate_store.IS_POSTGRES, climate_store._PG_DSN, climate_store._pool_obj) climate_store.IS_POSTGRES = True climate_store._PG_DSN = climate_store._libpq_dsn(url) - if hasattr(climate_store._local, "conn"): - del climate_store._local.conn + climate_store._pool_obj = None # force a fresh pool bound to the test DSN ddl_cols = ", ".join(f"{c} DOUBLE PRECISION" for c in climate_store.COLS if c != "date") with psycopg.connect(climate_store._PG_DSN, autocommit=True) as conn: @@ -100,9 +99,9 @@ def pg_store(): yield climate_store - climate_store.IS_POSTGRES, climate_store._PG_DSN = saved - if hasattr(climate_store._local, "conn"): - del climate_store._local.conn + if climate_store._pool_obj is not None: + climate_store._pool_obj.close() + climate_store.IS_POSTGRES, climate_store._PG_DSN, climate_store._pool_obj = saved def test_pg_history_roundtrip_and_dtypes(pg_store): diff --git a/tests/data/test_store.py b/tests/data/test_store.py index 6cc5f47..fb0b6ac 100644 --- a/tests/data/test_store.py +++ b/tests/data/test_store.py @@ -1,7 +1,12 @@ +"""data/store.py: hermetic SQLite tests (always run, via the tmp_store fixture in +conftest.py) plus a gated Postgres round-trip through the bounded connection pool +(skipped unless THERMOGRAPH_TEST_DATABASE_URL points at a reachable Postgres).""" import json import sqlite3 import time +import pytest + def test_payload_round_trip(tmp_store): payload = {"cell": "1_2", "days": [1, 2, 3]} @@ -63,3 +68,86 @@ def test_store_is_a_pure_accelerator_when_db_unavailable(tmp_store, monkeypatch) assert json.loads(body) == {"v": 1} # still serves the encoded payload assert tmp_store.get_revgeo("x") == (False, None) tmp_store.put_revgeo("x", "label") # no-op, no exception + + +# --- gated Postgres integration (the bounded pool, _get_pg_pool) ----------- + +@pytest.fixture +def pg_store(): + """Point store.py at THERMOGRAPH_TEST_DATABASE_URL for one test -- exercises + the real bounded ConnectionPool (borrow-per-operation, configure=_configure_pg + running the schema DDL) rather than the SQLite fallback. + + store.py picks its dialect-specific SQL text (``?`` vs ``%s`` placeholders, + INSERT OR REPLACE vs ON CONFLICT) once, at import time, from IS_POSTGRES -- + correct in production (THERMOGRAPH_DATABASE_URL is set before the module is + ever imported) but not something flipping IS_POSTGRES back on afterwards + retroactively fixes, so this fixture swaps the SQL constants too, not just + IS_POSTGRES/_PG_DSN. Restores all of it afterwards.""" + import os + + from data import store + + url = os.environ.get("THERMOGRAPH_TEST_DATABASE_URL", "").strip() + if not url: + pytest.skip("set THERMOGRAPH_TEST_DATABASE_URL to run the Postgres pool test") + + saved = (store.IS_POSTGRES, store._PG_DSN, store._pg_pool, + store._SQL_GET_PAYLOAD, store._SQL_PUT_PAYLOAD, + store._SQL_GET_REVGEO, store._SQL_PUT_REVGEO) + store.IS_POSTGRES = True + store._PG_DSN = store._libpq_dsn(url) + store._pg_pool = None # force a fresh pool bound to the test DSN + store._SQL_GET_PAYLOAD = "SELECT token, payload FROM derived WHERE kind=%s AND cell_id=%s AND key=%s" + store._SQL_PUT_PAYLOAD = ( + "INSERT INTO derived (kind, cell_id, key, token, payload, updated_at) " + "VALUES (%s,%s,%s,%s,%s,%s) " + "ON CONFLICT (kind, cell_id, key) DO UPDATE SET " + "token=EXCLUDED.token, payload=EXCLUDED.payload, updated_at=EXCLUDED.updated_at" + ) + store._SQL_GET_REVGEO = "SELECT label, updated_at FROM revgeo WHERE key=%s" + store._SQL_PUT_REVGEO = ( + "INSERT INTO revgeo (key, label, updated_at) VALUES (%s,%s,%s) " + "ON CONFLICT (key) DO UPDATE SET label=EXCLUDED.label, updated_at=EXCLUDED.updated_at" + ) + + import psycopg + with psycopg.connect(store._PG_DSN, autocommit=True) as conn: + conn.execute("DROP TABLE IF EXISTS derived") + conn.execute("DROP TABLE IF EXISTS revgeo") + + yield store + + if store._pg_pool is not None: + store._pg_pool.close() + (store.IS_POSTGRES, store._PG_DSN, store._pg_pool, + store._SQL_GET_PAYLOAD, store._SQL_PUT_PAYLOAD, + store._SQL_GET_REVGEO, store._SQL_PUT_REVGEO) = saved + + +def test_pg_payload_round_trip_through_the_pool(pg_store): + """Two consecutive calls borrow (and return) separate pooled connections -- + unlike the old thread-local model there is no single connection object tying + the write and the read together, so this also exercises that the schema + (applied via the pool's configure= callback) is visible across borrows.""" + body = pg_store.put_payload("grade", "1_2", "k", "tok", {"v": 1}) + assert json.loads(body) == {"v": 1} + assert pg_store.get_payload("grade", "1_2", "k", "tok") == body + assert pg_store.get_payload("grade", "1_2", "k", "tok-mismatch") is None + + +def test_pg_revgeo_round_trip_through_the_pool(pg_store): + key = pg_store.revgeo_key(47.6, -122.3) + assert pg_store.get_revgeo(key) == (False, None) + pg_store.put_revgeo(key, "Somewhere, WA") + assert pg_store.get_revgeo(key) == (True, "Somewhere, WA") + + +def test_pg_pool_survives_many_sequential_borrows(pg_store): + """A crude stand-in for the finding this pool exists to fix: many operations + in a row must keep working off a small bounded pool (max_size=8, well below + this loop) instead of piling up unclosed connections.""" + for i in range(25): + pg_store.put_payload("grade", "cell", f"k{i}", "tok", {"i": i}) + for i in range(25): + assert pg_store.get_json("grade", "cell", f"k{i}", "tok") == {"i": i} diff --git a/tests/test_warm_cities.py b/tests/test_warm_cities.py new file mode 100644 index 0000000..71ad619 --- /dev/null +++ b/tests/test_warm_cities.py @@ -0,0 +1,38 @@ +"""warm_cities.py's cross-process single-instance guard (LOCK_PATH + the +core/singleton.py flock it reuses) -- an overlapping deploy launching a second +copy while the first is still warming must stand down, not double-spend +archive-fetch quota.""" +import fcntl +import os +import subprocess +import sys + +import warm_cities + + +def test_lock_path_lives_under_data_dir(): + assert os.path.dirname(warm_cities.LOCK_PATH) == warm_cities.paths.DATA_DIR + assert os.path.basename(warm_cities.LOCK_PATH) == "warm_cities.lock" + + +def test_second_instance_stands_down_and_exits_zero(tmp_path): + """With the lock already held (simulating an overlapping deploy invocation), a + second run of the script must log and exit 0 immediately -- never reaching + main()'s network/parquet-write path.""" + lock = str(tmp_path / "warm_cities.lock") + fd = os.open(lock, os.O_RDWR | os.O_CREAT, 0o644) + fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB) + try: + env = dict(os.environ, THERMOGRAPH_DATA_DIR=str(tmp_path)) + # If the guard didn't work, this would try a real network fetch (or hang) + # instead of returning almost instantly with exit code 0. + result = subprocess.run( + [sys.executable, warm_cities.__file__, "--limit", "0"], + cwd=os.path.dirname(os.path.abspath(warm_cities.__file__)), + env=env, capture_output=True, text=True, timeout=20, + ) + assert result.returncode == 0 + assert "already running" in result.stdout + finally: + fcntl.flock(fd, fcntl.LOCK_UN) + os.close(fd) diff --git a/tests/web/test_gzip.py b/tests/web/test_gzip.py new file mode 100644 index 0000000..13bfae7 --- /dev/null +++ b/tests/web/test_gzip.py @@ -0,0 +1,12 @@ +"""GZipMiddleware's compresslevel: it runs synchronously on the event loop, so a +lower level (6, the standard speed/ratio sweet spot) trades a little ratio for +noticeably less per-request CPU competing with every other async route.""" +from starlette.middleware.gzip import GZipMiddleware + +from web import app as appmod + + +def test_gzip_middleware_uses_a_lowered_compresslevel(): + mw = next(m for m in appmod.app.user_middleware if m.cls is GZipMiddleware) + assert mw.kwargs.get("compresslevel") == 6 + assert mw.kwargs.get("compresslevel") != 9 # not Starlette's CPU-heavy default diff --git a/warm_cities.py b/warm_cities.py index 76f39f3..199a726 100644 --- a/warm_cities.py +++ b/warm_cities.py @@ -8,14 +8,29 @@ Idempotent: a cell whose archive is already cached is skipped. Fetches are paced (default 2s) to stay well under the archive API's rate limit. A cell that still has no cached archive when its page is first requested self-heals via get_history, so this is an optimization, not a hard dependency. + +Deploy launches this script detached on every deploy with no overlap check, so a +slow previous run (still mid-warm) can still be going when the next deploy's +invocation starts. Two overlapping runs would double-spend archive-fetch quota on +the same cells and race climate._write_cache's parquet writes, so the CLI entry +point below claims a single-instance flock (see LOCK_PATH) before calling main(). """ +import os import sys import time from api import homepage +from core import singleton from data import cities from data import climate from data import grid +import paths + +# core/singleton.py's flock guard, reused here for cross-*process* (not +# cross-worker) exclusion: the first invocation holds this for its process +# lifetime; a second one (an overlapping deploy) fails the non-blocking flock +# and stands down immediately. +LOCK_PATH = os.path.join(paths.DATA_DIR, "warm_cities.lock") def main(limit: int | None = None, pace: float = 2.0) -> None: @@ -53,6 +68,12 @@ def main(limit: int | None = None, pace: float = 2.0) -> None: if __name__ == "__main__": + if not singleton.claim(LOCK_PATH): + # Not an error: an overlapping run just means a previous deploy's warm is + # still in flight. Exit 0 so the deploy script doesn't treat this as a + # failure. + print("warm_cities: another instance is already running -- exiting") + sys.exit(0) args = sys.argv[1:] lim = int(args[args.index("--limit") + 1]) if "--limit" in args else None pc = float(args[args.index("--pace") + 1]) if "--pace" in args else 2.0 diff --git a/web/app.py b/web/app.py index 0de9cf8..59217ae 100644 --- a/web/app.py +++ b/web/app.py @@ -181,7 +181,16 @@ app = FastAPI(title="Thermograph", version="0.2.0", lifespan=_lifespan) # Compress every sizeable response (the 2-year calendar JSON shrinks ~6-8×). # Applies to API JSON and static assets alike; tiny responses are left alone. -app.add_middleware(GZipMiddleware, minimum_size=1024) +# compresslevel: GZipMiddleware compresses synchronously ON THE EVENT LOOP, so +# this is CPU work competing with every other async route in the worker while it +# runs — zlib's default level 9 spends a lot of that time chasing the last couple +# percent of ratio for a payload this size. 6 is the standard speed/ratio sweet +# spot (roughly zlib's own recommended default): most of level 9's compression +# but a fraction of the CPU, so a big calendar/day/cell-bundle response doesn't +# stall other requests on this worker for as long. Not worth a custom threaded +# compression middleware for payloads this size — the level drop captures most +# of the win for near-zero complexity. +app.add_middleware(GZipMiddleware, minimum_size=1024, compresslevel=6) # CORS (repo-split Stage 5): unset (default) means no CORSMiddleware at all -- # today's exact behavior, browsers already allow same-origin fetch with no From 83c2e05b967d127f37f151a2d82f7be9a79c54b4 Mon Sep 17 00:00:00 2001 From: emi Date: Thu, 23 Jul 2026 04:41:32 +0000 Subject: [PATCH 3/3] Land the Discord gateway bot (port + hardening) (#7) --- notifications/discord_bot.py | 295 ++++++++++++++++++++++++ requirements.txt | 3 + tests/notifications/test_discord_bot.py | 125 ++++++++++ web/app.py | 17 ++ 4 files changed, 440 insertions(+) create mode 100644 notifications/discord_bot.py create mode 100644 tests/notifications/test_discord_bot.py diff --git a/notifications/discord_bot.py b/notifications/discord_bot.py new file mode 100644 index 0000000..6e32780 --- /dev/null +++ b/notifications/discord_bot.py @@ -0,0 +1,295 @@ +"""Discord **gateway** bot — reacts to @mentions (and DMs) in real time. + +The other two Discord paths need no persistent connection: slash commands arrive as +signed HTTP POSTs (discord_interactions.py) and the daily post goes out over an +incoming webhook (discord.py). This one is different — to answer an *ordinary* +message that mentions the bot, we have to hold a live websocket to Discord's +gateway and listen for MESSAGE_CREATE events. + +Single-connection invariant. Discord permits exactly one gateway connection per bot +token (per shard), so — like the subscription notifier and the daily post — this +runs in exactly one process: web/app.py starts it only when the process wins the +leader election (see core/singleton.claim_leader). No new container or daemon; it +rides the app's asyncio event loop as a background task. + +Least privilege / no privileged intent. We subscribe to GUILD_MESSAGES + +DIRECT_MESSAGES and act only on messages that @mention the bot or DM it. Discord +delivers message *content* for exactly those cases even without the privileged +MESSAGE_CONTENT intent, so the bot works with no portal review. To trigger on +arbitrary keywords in channels, enable the MESSAGE_CONTENT intent in the Developer +Portal and extend `_response_for_message` — the rest of the machinery is unchanged. + +Enable with THERMOGRAPH_DISCORD_BOT=1 (plus THERMOGRAPH_DISCORD_BOT_TOKEN, the same +bot token discord.py uses for DMs). Unset => no gateway connection; everything here +no-ops, so an unconfigured deploy pays nothing. + +Behaviour: mention the bot with a city name ("@Thermograph Phoenix") and it replies +with the same grade embed the /grade slash command returns, reusing +discord_interactions._grade_message so the two never drift. +""" +from __future__ import annotations + +import asyncio +import json +import logging +import os +import random + +from notifications import discord + +log = logging.getLogger("thermograph.discord_bot") + +# v10 JSON gateway. On a resume we reconnect to the session's resume_gateway_url +# (from the READY payload) instead, with the same query string appended. +_GATEWAY_URL = "wss://gateway.discord.gg/?v=10&encoding=json" +_GATEWAY_QS = "/?v=10&encoding=json" + +# Gateway intents (bitfield). GUILDS gives guild/channel context; GUILD_MESSAGES and +# DIRECT_MESSAGES deliver message-create events in servers and DMs. We deliberately +# do NOT request MESSAGE_CONTENT (privileged) — Discord still includes content for +# messages that mention us or DM us, which is all we act on. +_INTENT_GUILDS = 1 << 0 +_INTENT_GUILD_MESSAGES = 1 << 9 +_INTENT_DIRECT_MESSAGES = 1 << 12 +INTENTS = _INTENT_GUILDS | _INTENT_GUILD_MESSAGES | _INTENT_DIRECT_MESSAGES + +# Gateway opcodes (discord.com/developers/docs/topics/gateway-events#opcodes). +_OP_DISPATCH = 0 +_OP_HEARTBEAT = 1 +_OP_IDENTIFY = 2 +_OP_RESUME = 6 +_OP_RECONNECT = 7 +_OP_INVALID_SESSION = 9 +_OP_HELLO = 10 +_OP_HEARTBEAT_ACK = 11 + +# Close codes for unrecoverable conditions — reconnecting just loops, so we stop: +# 4004 auth failed (bad token), 4010 invalid shard, 4011 sharding required, 4012 +# invalid API version, 4013 invalid intents, 4014 disallowed (privileged) intents. +_FATAL_CLOSE = {4004, 4010, 4011, 4012, 4013, 4014} + +_TRUTHY = {"1", "true", "yes", "on"} + +_HELP = ("Mention me with a city name — e.g. `@Thermograph Phoenix` — and I'll grade " + "today's weather against ~45 years of local history.") + +# Reconnect backoff bounds (seconds). +_BACKOFF_START = 1.0 +_BACKOFF_MAX = 60.0 + + +def enabled() -> bool: + """True iff the gateway bot is switched on AND a bot token is configured. Off by + default: a gateway connection is a new persistent behaviour, so it takes an + explicit opt-in rather than piggybacking on the token being present for DMs.""" + flag = os.environ.get("THERMOGRAPH_DISCORD_BOT", "").strip().lower() in _TRUTHY + return flag and bool(discord.BOT_TOKEN) + + +# --- message handling (pure; unit-tested without a live gateway) -------------- +def _mentions_bot(msg: dict, bot_user_id: str) -> bool: + return any(str(u.get("id")) == str(bot_user_id) for u in (msg.get("mentions") or [])) + + +def _is_dm(msg: dict) -> bool: + # A DM has no guild_id; a guild message always carries one. + return not msg.get("guild_id") + + +def _strip_mentions(content: str, bot_user_id: str) -> str: + """Drop every form of the bot's own mention (<@id> and the legacy <@!id>) and + collapse the surrounding whitespace, leaving just the user's query text.""" + out = content.replace(f"<@{bot_user_id}>", " ").replace(f"<@!{bot_user_id}>", " ") + return " ".join(out.split()) + + +def _response_for_message(msg: dict, bot_user_id: str) -> dict | None: + """The reply payload for one MESSAGE_CREATE, or None to stay silent. + + Silent unless the message DMs the bot or @mentions it, and never for a bot + author (that includes ourselves — the guard against a mention-loop). The text + after the mention is treated as a city query and answered with the same embed + the /grade slash command builds; an empty query gets the help line.""" + author = msg.get("author") or {} + if author.get("bot") or str(author.get("id")) == str(bot_user_id): + return None + dm = _is_dm(msg) + if not dm and not _mentions_bot(msg, bot_user_id): + return None + content = msg.get("content") or "" + query = content.strip() if dm else _strip_mentions(content, bot_user_id) + if not query: + return {"content": _HELP} + # Import here (not at module top) to avoid importing the FastAPI web layer's + # transitive deps for a module that also loads in the notifier-only context. + from notifications import discord_interactions + data = discord_interactions._grade_message(query) + # Gateway messages can't be ephemeral — the flag is interactions-only. + data.pop("flags", None) + return data + + +# --- REST reply --------------------------------------------------------------- +async def _send_reply(channel_id, message_id, data: dict) -> None: + """Post the reply in-channel as a proper reply to the triggering message. + + Reuses discord._bot_post (sync httpx with a 429 retry) off the event loop via a + thread so the gateway read-loop never blocks. allowed_mentions is locked down so + a crafted query can't turn the reply into an @everyone/role ping; only the person + who asked is pinged, via the reply reference.""" + if not channel_id: + return + payload = dict(data) + if message_id: + payload["message_reference"] = {"message_id": str(message_id)} + payload["allowed_mentions"] = {"parse": [], "replied_user": True} + try: + await asyncio.to_thread( + discord._bot_post, f"/channels/{channel_id}/messages", payload) + except Exception: # noqa: BLE001 - a failed reply must never kill the read loop + log.warning("discord_bot: reply failed", exc_info=True) + + +# --- gateway session ---------------------------------------------------------- +class _Session: + """Mutable per-connection state that survives a resume: the last sequence + number, the resume token/URL, and our own user id (for mention detection).""" + + __slots__ = ("seq", "session_id", "resume_url", "user_id") + + def __init__(self) -> None: + self.seq: int | None = None + self.session_id: str | None = None + self.resume_url: str | None = None + self.user_id: str = "" + + +async def _heartbeat(ws, interval: float, sess: _Session, acked: dict) -> None: + """Send op-1 heartbeats every `interval` seconds (first beat jittered per + Discord's guidance). If the previous beat was never ACKed the link is a zombie — + close with a non-1000 code so the next connect can resume, and stop.""" + await asyncio.sleep(interval * random.random()) + while True: + if not acked["ok"]: + await ws.close(code=4000) + return + acked["ok"] = False + try: + await ws.send(json.dumps({"op": _OP_HEARTBEAT, "d": sess.seq})) + except Exception: # noqa: BLE001 - closed socket; the read loop handles reconnect + return + await asyncio.sleep(interval) + + +async def _dispatch(msg: dict, sess: _Session) -> None: + """Handle an op-0 DISPATCH event we care about.""" + t = msg.get("t") + d = msg.get("d") or {} + if t == "READY": + sess.session_id = d.get("session_id") + resume = (d.get("resume_gateway_url") or "").rstrip("/") + sess.resume_url = f"{resume}{_GATEWAY_QS}" if resume else None + sess.user_id = str((d.get("user") or {}).get("id") or "") + log.info("discord_bot: gateway ready (user %s)", sess.user_id) + elif t == "MESSAGE_CREATE" and sess.user_id: + # _response_for_message reaches into discord_interactions._grade_message for + # anything but the help line -- sync parquet/DB reads + grading math, same as + # web/app.py's discord_interactions route -- keep it off the (single, shared) + # gateway event loop so one grading lookup can't stall heartbeats/reads. + reply = await asyncio.to_thread(_response_for_message, d, sess.user_id) + if reply: + await _send_reply(d.get("channel_id"), d.get("id"), reply) + + +async def _connection(ws, sess: _Session, resume: bool) -> str: + """Drive one live connection until it closes. Returns the mode for the next + connect: "resume" to reconnect and RESUME, "identify" to start a fresh session.""" + hello = json.loads(await ws.recv()) + if hello.get("op") != _OP_HELLO: + return "identify" # protocol out of step — start clean + interval = float(hello["d"]["heartbeat_interval"]) / 1000.0 + acked = {"ok": True} + hb = asyncio.create_task(_heartbeat(ws, interval, sess, acked)) + try: + if resume and sess.session_id and sess.seq is not None: + await ws.send(json.dumps({"op": _OP_RESUME, "d": { + "token": discord.BOT_TOKEN, + "session_id": sess.session_id, + "seq": sess.seq, + }})) + else: + await ws.send(json.dumps({"op": _OP_IDENTIFY, "d": { + "token": discord.BOT_TOKEN, + "intents": INTENTS, + "properties": {"os": "linux", "browser": "thermograph", "device": "thermograph"}, + }})) + async for raw in ws: + msg = json.loads(raw) + if msg.get("s") is not None: + sess.seq = msg["s"] + op = msg.get("op") + if op == _OP_DISPATCH: + await _dispatch(msg, sess) + elif op == _OP_HEARTBEAT: # server asked for an immediate beat + await ws.send(json.dumps({"op": _OP_HEARTBEAT, "d": sess.seq})) + elif op == _OP_HEARTBEAT_ACK: + acked["ok"] = True + elif op == _OP_RECONNECT: # planned reconnect — resume + return "resume" + elif op == _OP_INVALID_SESSION: + # d == True => the session can still be resumed; else it's dead and + # the next connect must IDENTIFY fresh. + if not msg.get("d"): + sess.session_id = None + sess.resume_url = None + await asyncio.sleep(1 + random.random()) + return "resume" if sess.session_id else "identify" + finally: + hb.cancel() + return "resume" if sess.session_id else "identify" + + +def _close_code(exc: Exception): + """Best-effort websocket close code across websockets library versions (older + exposes .code; newer nests it under .rcvd).""" + code = getattr(exc, "code", None) + if code: + return code + return getattr(getattr(exc, "rcvd", None), "code", None) + + +async def run() -> None: + """Maintain the gateway connection for the app's lifetime: connect, IDENTIFY (or + RESUME), dispatch events, and reconnect with capped exponential backoff. Returns + only on a fatal close (bad token / disallowed intents) or task cancellation, and + never raises out — a misbehaving gateway must not take the app down.""" + if not enabled(): + return + try: + import websockets + except Exception: # noqa: BLE001 - shipped via uvicorn[standard]; guard anyway + log.warning("discord_bot: 'websockets' unavailable; gateway bot disabled") + return + sess = _Session() + mode = "identify" + backoff = _BACKOFF_START + while True: + use_resume = mode == "resume" and bool(sess.resume_url) + url = sess.resume_url if use_resume else _GATEWAY_URL + try: + async with websockets.connect(url, max_size=2 ** 20, open_timeout=30) as ws: + mode = await _connection(ws, sess, resume=use_resume) + backoff = _BACKOFF_START # server-requested reconnect: come back promptly + continue + except asyncio.CancelledError: + log.info("discord_bot: shutting down") + raise + except Exception as exc: # noqa: BLE001 - any drop -> back off and reconnect + if _close_code(exc) in _FATAL_CLOSE: + log.error("discord_bot: fatal gateway close %s; stopping", _close_code(exc)) + return + mode = "resume" if sess.session_id else "identify" + log.warning("discord_bot: gateway dropped (%s); reconnecting in %.0fs", + type(exc).__name__, backoff) + await asyncio.sleep(backoff + random.random()) + backoff = min(backoff * 2, _BACKOFF_MAX) diff --git a/requirements.txt b/requirements.txt index c12d96d..62fb0fc 100644 --- a/requirements.txt +++ b/requirements.txt @@ -19,6 +19,9 @@ alembic==1.14.0 pywebpush==2.0.0 # Ed25519 verification of Discord interaction webhooks (see notifications/discord_interactions.py). PyNaCl==1.5.0 +# Discord gateway bot's websocket client (notifications/discord_bot.py). Already +# pulled in transitively by uvicorn[standard]; pinned here since we import it directly. +websockets==16.0 # Recurring worker-tier jobs — city warming, IndexNow pings (see # notifications/scheduler.py). In-process; no broker/queue needed at this scale. apscheduler==3.11.3 diff --git a/tests/notifications/test_discord_bot.py b/tests/notifications/test_discord_bot.py new file mode 100644 index 0000000..9692542 --- /dev/null +++ b/tests/notifications/test_discord_bot.py @@ -0,0 +1,125 @@ +"""Discord gateway bot: the message-handling logic that decides when to reply and +with what. No live gateway — we feed MESSAGE_CREATE payloads straight to the pure +handler and stub the shared /grade lookup so these stay data-independent.""" +import asyncio + +import pytest + +from notifications import discord +from notifications import discord_bot as bot +from notifications import discord_interactions as di + +BOT_ID = "999" + + +@pytest.fixture(autouse=True) +def _stub_grade(monkeypatch): + """Reply payload for /grade, mirroring the real shape (embeds + an ephemeral + flag the gateway path must drop).""" + monkeypatch.setattr( + di, "_grade_message", + lambda query: {"embeds": [{"title": f"grade:{query}"}], "flags": 64}) + + +def _msg(content="", *, author_id="1", bot_author=False, mentions=(), guild=True): + m = { + "content": content, + "id": "42", + "channel_id": "77", + "author": {"id": author_id, "bot": bot_author}, + "mentions": [{"id": mid} for mid in mentions], + } + if guild: + m["guild_id"] = "5" + return m + + +# --- when the bot stays silent ---------------------------------------------- + +def test_ignores_other_bots(): + assert bot._response_for_message( + _msg("<@999> Phoenix", author_id="8", bot_author=True, mentions=[BOT_ID]), BOT_ID) is None + + +def test_ignores_its_own_messages(): + assert bot._response_for_message( + _msg("hello", author_id=BOT_ID, mentions=[BOT_ID]), BOT_ID) is None + + +def test_ignores_guild_message_without_a_mention(): + assert bot._response_for_message(_msg("just chatting", mentions=[]), BOT_ID) is None + + +# --- when the bot replies ---------------------------------------------------- + +def test_mention_with_city_returns_grade_without_ephemeral_flag(): + out = bot._response_for_message(_msg("<@999> Phoenix", mentions=[BOT_ID]), BOT_ID) + assert out == {"embeds": [{"title": "grade:Phoenix"}]} # flags dropped + assert "flags" not in out + + +def test_legacy_nickname_mention_is_stripped(): + out = bot._response_for_message(_msg("<@!999> Tokyo", mentions=[BOT_ID]), BOT_ID) + assert out["embeds"][0]["title"] == "grade:Tokyo" + + +def test_dm_needs_no_mention(): + out = bot._response_for_message(_msg("Berlin", guild=False, mentions=[]), BOT_ID) + assert out["embeds"][0]["title"] == "grade:Berlin" + + +def test_mention_with_no_query_gives_help(): + out = bot._response_for_message(_msg("<@999>", mentions=[BOT_ID]), BOT_ID) + assert "content" in out and "mention me" in out["content"].lower() + + +# --- helpers ----------------------------------------------------------------- + +def test_strip_mentions_handles_both_forms(): + assert bot._strip_mentions("<@999> <@!999> Sao Paulo", BOT_ID) == "Sao Paulo" + + +def test_is_dm_is_true_without_guild_id(): + assert bot._is_dm(_msg("x", guild=False)) is True + assert bot._is_dm(_msg("x", guild=True)) is False + + +# --- _dispatch: same reply, off the event loop ------------------------------- + +def test_dispatch_delivers_the_same_reply_via_a_thread(monkeypatch): + """_dispatch now runs _response_for_message through asyncio.to_thread (the + event-loop fix) -- confirm that still produces the exact same reply and gets + it to _send_reply, i.e. the offload didn't change behaviour.""" + sent = {} + + async def _fake_send_reply(channel_id, message_id, data): + sent["channel_id"] = channel_id + sent["message_id"] = message_id + sent["data"] = data + + monkeypatch.setattr(bot, "_send_reply", _fake_send_reply) + sess = bot._Session() + sess.user_id = BOT_ID + msg = {"t": "MESSAGE_CREATE", "d": _msg("<@999> Phoenix", mentions=[BOT_ID])} + + asyncio.run(bot._dispatch(msg, sess)) + + assert sent["channel_id"] == "77" + assert sent["message_id"] == "42" + assert sent["data"] == {"embeds": [{"title": "grade:Phoenix"}]} + + +# --- enabled() gating -------------------------------------------------------- + +def test_enabled_is_off_by_default(monkeypatch): + monkeypatch.delenv("THERMOGRAPH_DISCORD_BOT", raising=False) + monkeypatch.setattr(discord, "BOT_TOKEN", "a-token") + assert bot.enabled() is False + + +def test_enabled_needs_both_flag_and_token(monkeypatch): + monkeypatch.setenv("THERMOGRAPH_DISCORD_BOT", "1") + monkeypatch.setattr(discord, "BOT_TOKEN", "") + assert bot.enabled() is False + monkeypatch.setattr(discord, "BOT_TOKEN", "a-token") + assert bot.enabled() is True diff --git a/web/app.py b/web/app.py index 59217ae..0d6d8bb 100644 --- a/web/app.py +++ b/web/app.py @@ -1,4 +1,5 @@ """Thermograph API — grade recent local weather against ~45 years of climatology.""" +import asyncio import contextlib import datetime import functools @@ -22,6 +23,7 @@ from api import content_routes from core import audit from data import climate from notifications import digest +from notifications import discord_bot from notifications import discord_interactions as discord_interactions_mod from notifications import discord_link from accounts import db @@ -162,16 +164,31 @@ async def _lifespan(app): # restricts this to processes that are allowed to own it at all — "web" replicas # never start it even if they'd otherwise win the leader election, so a stateless # N-replica web tier can scale without also scaling notifier instances. + bot_task = None if _should_run_notifier(): notify.start() # Same leader gate as the notifier above — the recurring worker jobs # (city warming, IndexNow) must likewise run in exactly one process. scheduler.start() + # Discord gateway bot (opt-in). Discord allows one gateway connection per + # bot token, so it belongs on the same single leader; it rides this event + # loop as a background task rather than a new process. No-op unless + # THERMOGRAPH_DISCORD_BOT is set with a bot token (discord_bot.enabled()). + if discord_bot.enabled(): + bot_task = asyncio.create_task(discord_bot.run(), name="discord-bot") # Deploy/restart marker: one line correlates a metric shift with a boot. audit.log_activity("app.start", {"version": app.version, "role": ROLE, "base": BASE, "notifier": _should_run_notifier(), "pid": os.getpid()}) yield audit.log_activity("app.stop", {"role": ROLE, "pid": os.getpid()}) + if bot_task is not None: + bot_task.cancel() + # Wait for the cancellation to actually land before tearing anything else + # down -- otherwise shutdown races the bot coroutine's own teardown (closing + # the websocket, cancelling its heartbeat task) and the CancelledError it + # raises on its way out surfaces as unhandled-task noise in the logs. + with contextlib.suppress(asyncio.CancelledError): + await bot_task notify.stop() scheduler.stop() await _frontend_client.aclose()