Merge pull request 'Make unread Web Push delivery multi-worker safe' (#552) from timmy/551-push-delivery-lease into main
All checks were successful
CI / lint (push) Successful in 1m5s
CI / build-release (push) Successful in 5s
CI / release-candidate (push) Successful in 6s

This commit is contained in:
timmy 2026-08-11 07:27:55 +00:00
commit 1f05521094
5 changed files with 229 additions and 33 deletions

View File

@ -192,8 +192,12 @@ export STACKCHAIN_TRUSTED_PROXY_CIDRS='127.0.0.0/8'
export STACKCHAIN_VAPID_PUBLIC_KEY='<url-safe-public-key>'
export STACKCHAIN_VAPID_PRIVATE_KEY='<private-key-from-secret-manager>'
export STACKCHAIN_VAPID_SUBJECT='mailto:ops@example.com'
# Optional; defaults to 30 seconds and STACKCHAIN_STATE_DIR/push-subscriptions.sqlite3.
# Optional; defaults to a 30-second poll, 10-second endpoint deadline,
# 60-second renewable cross-worker lease, and
# STACKCHAIN_STATE_DIR/push-subscriptions.sqlite3.
export STACKCHAIN_PUSH_POLL_SECONDS=30
export STACKCHAIN_PUSH_SEND_TIMEOUT_SECONDS=10
export STACKCHAIN_PUSH_LEASE_SECONDS=60
export STACKCHAIN_PUSH_DB='/var/lib/stackchain-dashboard/push-subscriptions.sqlite3'
uvicorn src.main:app --host 127.0.0.1 --port 8000
```

View File

@ -80,6 +80,13 @@ async def _drain_authored_action_operations() -> None:
async def _push_poll_loop() -> None:
interval = max(5.0, float(os.getenv("STACKCHAIN_PUSH_POLL_SECONDS", "30")))
send_timeout = max(
1.0, float(os.getenv("STACKCHAIN_PUSH_SEND_TIMEOUT_SECONDS", "10"))
)
lease_seconds = max(
send_timeout + 5.0,
float(os.getenv("STACKCHAIN_PUSH_LEASE_SECONDS", "60")),
)
while True:
await asyncio.sleep(interval)
try:
@ -87,6 +94,8 @@ async def _push_poll_loop() -> None:
_push_subscription_store,
_push_configuration(),
notifications,
lease_seconds=lease_seconds,
send_timeout_seconds=send_timeout,
)
except asyncio.CancelledError:
raise

View File

@ -1,5 +1,7 @@
import asyncio
import json
import secrets
import time
from dataclasses import dataclass
from typing import Awaitable, Callable
@ -37,38 +39,68 @@ async def dispatch_unread_updates(
configuration: PushConfiguration,
unread: Callable[[], Awaitable[dict]],
send: Callable[[dict, str], Awaitable[None]] | None = None,
*,
lease_seconds: float = 60.0,
send_timeout_seconds: float = 10.0,
) -> int:
if not configuration.enabled:
return 0
page = await unread()
thread_ids = {
int(item["notification_id"])
for item in page.get("items", [])
if isinstance(item, dict) and str(item.get("notification_id", "")).isdigit()
}
count = 0
for delivery in store.claim_unseen(thread_ids):
for thread_id in delivery.thread_ids:
payload = json.dumps(
{
"title": "New work update",
"body": "Tap to review it in Stackchain.",
"route": f"#/my-work/update/{thread_id}",
"tag": f"stackchain-update-{thread_id}",
},
separators=(",", ":"),
)
try:
if send is None:
await send_web_push(delivery.subscription, payload, configuration)
else:
await send(delivery.subscription, payload)
except Exception as error:
status = getattr(getattr(error, "response", None), "status_code", None)
if status in {404, 410}:
store.delete_session(delivery.session_id)
break
raise
store.mark_delivered(delivery.session_id, (thread_id,))
count += 1
return count
owner = secrets.token_urlsafe(18)
acquired = await asyncio.to_thread(
store.acquire_dispatch_lease,
owner,
now=time.time(),
lease_seconds=lease_seconds,
)
if not acquired:
return 0
try:
page = await unread()
thread_ids = {
int(item["notification_id"])
for item in page.get("items", [])
if isinstance(item, dict) and str(item.get("notification_id", "")).isdigit()
}
count = 0
for delivery in await asyncio.to_thread(store.claim_unseen, thread_ids):
for thread_id in delivery.thread_ids:
still_owner = await asyncio.to_thread(
store.acquire_dispatch_lease,
owner,
now=time.time(),
lease_seconds=lease_seconds,
)
if not still_owner:
return count
payload = json.dumps(
{
"title": "New work update",
"body": "Tap to review it in Stackchain.",
"route": f"#/my-work/update/{thread_id}",
"tag": f"stackchain-update-{thread_id}",
},
separators=(",", ":"),
)
try:
if send is None:
operation = send_web_push(
delivery.subscription, payload, configuration
)
else:
operation = send(delivery.subscription, payload)
await asyncio.wait_for(operation, timeout=send_timeout_seconds)
except Exception as error:
status = getattr(getattr(error, "response", None), "status_code", None)
if status in {404, 410}:
await asyncio.to_thread(store.delete_session, delivery.session_id)
break
# A transient provider failure belongs to this endpoint;
# leave it unseen for a later poll and continue fan-out.
continue
await asyncio.to_thread(
store.mark_delivered, delivery.session_id, (thread_id,)
)
count += 1
return count
finally:
await asyncio.to_thread(store.release_dispatch_lease, owner)

View File

@ -1,4 +1,5 @@
import json
import os
import sqlite3
from dataclasses import dataclass
from pathlib import Path
@ -18,6 +19,7 @@ class PushSubscriptionStore:
def __init__(self, path: str | Path):
self.path = Path(path)
self.path.parent.mkdir(parents=True, exist_ok=True)
os.chmod(self.path.parent, 0o700)
with self._connect() as connection:
connection.executescript(
"""
@ -33,14 +35,47 @@ class PushSubscriptionStore:
FOREIGN KEY (session_id) REFERENCES push_subscriptions(session_id)
ON DELETE CASCADE
);
CREATE TABLE IF NOT EXISTS push_dispatch_lease (
singleton INTEGER PRIMARY KEY CHECK (singleton = 1),
owner TEXT NOT NULL,
expires_at REAL NOT NULL
);
"""
)
os.chmod(self.path, 0o600)
def _connect(self):
connection = sqlite3.connect(self.path, timeout=2)
connection.execute("PRAGMA foreign_keys = ON")
return connection
def acquire_dispatch_lease(
self, owner: str, *, now: float, lease_seconds: float
) -> bool:
with self._connect() as connection:
connection.execute("BEGIN IMMEDIATE")
current = connection.execute(
"SELECT owner, expires_at FROM push_dispatch_lease WHERE singleton = 1"
).fetchone()
if current is not None and current[0] != owner and current[1] > now:
return False
connection.execute(
"""INSERT INTO push_dispatch_lease(singleton, owner, expires_at)
VALUES (1, ?, ?)
ON CONFLICT(singleton) DO UPDATE SET
owner = excluded.owner, expires_at = excluded.expires_at""",
(owner, now + lease_seconds),
)
return True
def release_dispatch_lease(self, owner: str) -> bool:
with self._connect() as connection:
result = connection.execute(
"DELETE FROM push_dispatch_lease WHERE singleton = 1 AND owner = ?",
(owner,),
)
return result.rowcount == 1
def upsert(self, session_id: str, subscription: dict) -> None:
endpoint = subscription["endpoint"]
encoded = json.dumps(subscription, separators=(",", ":"), sort_keys=True)

View File

@ -1,4 +1,6 @@
import json
import asyncio
import os
from types import SimpleNamespace
import pytest
@ -8,6 +10,32 @@ from src.push_notifications import PushConfiguration, dispatch_unread_updates
from src.push_subscription_store import PushSubscriptionStore
def test_dispatch_lease_is_exclusive_recoverable_and_owner_fenced(tmp_path):
path = tmp_path / "push.sqlite3"
first = PushSubscriptionStore(path)
second = PushSubscriptionStore(path)
assert first.acquire_dispatch_lease("worker-a", now=100, lease_seconds=30) is True
assert second.acquire_dispatch_lease("worker-b", now=100, lease_seconds=30) is False
assert second.acquire_dispatch_lease("worker-b", now=131, lease_seconds=30) is True
assert first.release_dispatch_lease("worker-a") is False
assert second.release_dispatch_lease("worker-b") is True
def test_subscription_store_uses_private_filesystem_permissions(tmp_path):
state_dir = tmp_path / "push-state"
previous_umask = os.umask(0)
try:
database = state_dir / "push.sqlite3"
PushSubscriptionStore(database)
finally:
os.umask(previous_umask)
assert state_dir.stat().st_mode & 0o777 == 0o700
assert database.stat().st_mode & 0o777 == 0o600
def test_subscription_store_delivers_each_unread_thread_once_per_device(tmp_path):
store = PushSubscriptionStore(tmp_path / "push.sqlite3")
subscription = {
@ -79,6 +107,40 @@ async def test_dispatch_sends_one_privacy_safe_deep_link_per_new_thread(tmp_path
assert "Secret title" not in json.dumps(sent)
@pytest.mark.anyio
async def test_concurrent_workers_do_not_dispatch_the_same_update(tmp_path):
path = tmp_path / "push.sqlite3"
first = PushSubscriptionStore(path)
second = PushSubscriptionStore(path)
first.upsert("session-a", {
"endpoint": "https://push.example/device-a",
"keys": {"p256dh": "public-key", "auth": "auth-secret"},
})
sending = asyncio.Event()
finish = asyncio.Event()
sent = []
async def unread():
return {"items": [{"notification_id": 42}]}
async def send(_subscription, payload):
sent.append(json.loads(payload)["tag"])
sending.set()
await finish.wait()
config = PushConfiguration("public", "private", "mailto:ops@example.com")
active = asyncio.create_task(dispatch_unread_updates(first, config, unread, send))
await sending.wait()
competing = await asyncio.wait_for(
dispatch_unread_updates(second, config, unread, send), timeout=0.1
)
finish.set()
assert competing == 0
assert await active == 1
assert sent == ["stackchain-update-42"]
@pytest.mark.anyio
async def test_dispatch_removes_an_expired_push_endpoint(tmp_path):
store = PushSubscriptionStore(tmp_path / "push.sqlite3")
@ -101,6 +163,60 @@ async def test_dispatch_removes_an_expired_push_endpoint(tmp_path):
assert store.is_subscribed("session-a") is False
@pytest.mark.anyio
async def test_transient_endpoint_failure_does_not_block_healthy_devices(tmp_path):
store = PushSubscriptionStore(tmp_path / "push.sqlite3")
for session_id in ("session-a", "session-b"):
store.upsert(session_id, {
"endpoint": f"https://push.example/{session_id}",
"keys": {"p256dh": "public-key", "auth": "auth-secret"},
})
sent = []
async def unread():
return {"items": [{"notification_id": 8}]}
async def send(subscription, _payload):
if subscription["endpoint"].endswith("session-a"):
raise RuntimeError("provider unavailable")
sent.append(subscription["endpoint"])
config = PushConfiguration("public", "private", "mailto:ops@example.com")
assert await dispatch_unread_updates(store, config, unread, send) == 1
assert sent == ["https://push.example/session-b"]
assert store.claim_unseen({8})[0].session_id == "session-a"
@pytest.mark.anyio
async def test_timed_out_endpoint_does_not_stall_push_fanout(tmp_path):
store = PushSubscriptionStore(tmp_path / "push.sqlite3")
for session_id in ("session-a", "session-b"):
store.upsert(session_id, {
"endpoint": f"https://push.example/{session_id}",
"keys": {"p256dh": "public-key", "auth": "auth-secret"},
})
sent = []
async def unread():
return {"items": [{"notification_id": 13}]}
async def send(subscription, _payload):
if subscription["endpoint"].endswith("session-a"):
await asyncio.Event().wait()
sent.append(subscription["endpoint"])
config = PushConfiguration("public", "private", "mailto:ops@example.com")
delivered = await asyncio.wait_for(
dispatch_unread_updates(
store, config, unread, send, send_timeout_seconds=0.01
),
timeout=0.2,
)
assert delivered == 1
assert sent == ["https://push.example/session-b"]
@pytest.mark.anyio
async def test_authenticated_device_can_subscribe_report_status_and_unsubscribe(tmp_path, monkeypatch):
store = PushSubscriptionStore(tmp_path / "push.sqlite3")