stackchain-dashboard/src/push_notifications.py
timmy 13e73bd342
All checks were successful
CI / lint (pull_request) Successful in 2m53s
CI / build-release (pull_request) Successful in 6s
CI / browser-journey (pull_request) Successful in 3m1s
CI / release-candidate (pull_request) Has been skipped
feat: verify per-device Push delivery health (Closes #1120)
2026-08-19 09:00:06 +00:00

521 lines
20 KiB
Python

import asyncio
import json
import secrets
import time
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Awaitable, Callable
from urllib.parse import urlsplit
from zoneinfo import ZoneInfo, ZoneInfoNotFoundError
import requests
from src.push_subscription_store import PushSubscriptionStore
from src.push_endpoint_policy import (
ResolvedPushEndpoint,
UnsafePushEndpoint,
resolve_public_push_endpoint,
validate_public_push_endpoint,
)
@dataclass(frozen=True)
class PushConfiguration:
public_key: str
private_key: str
subject: str
@property
def enabled(self) -> bool:
return bool(self.public_key and self.private_key and self.subject)
class _PinnedHTTPSAdapter(requests.adapters.HTTPAdapter):
"""Dial one approved IP while authenticating the endpoint's original host."""
def __init__(self, resolved: ResolvedPushEndpoint):
self.resolved = resolved
super().__init__()
def add_headers(self, request, **kwargs):
super().add_headers(request, **kwargs)
request.headers["Host"] = self.resolved.hostname
def build_connection_pool_key_attributes(self, request, verify, cert=None):
parsed = urlsplit(request.url)
if parsed.scheme != "https" or parsed.hostname != self.resolved.hostname:
raise UnsafePushEndpoint("Push transport attempted an unvalidated destination")
host, tls = super().build_connection_pool_key_attributes(request, verify, cert)
host.update(
host=self.resolved.addresses[0],
port=self.resolved.port,
)
tls.update(
assert_hostname=self.resolved.hostname,
server_hostname=self.resolved.hostname,
)
return host, tls
def _delivery_failure_reason(error: Exception) -> str:
if isinstance(error, (asyncio.TimeoutError, TimeoutError)):
return "timeout"
return "provider"
async def send_web_push(
subscription: dict,
payload: str,
configuration: PushConfiguration,
*,
endpoint_resolver: Callable[[str], Awaitable[ResolvedPushEndpoint]] | None = None,
webpush_sender: Callable[..., object] | None = None,
) -> None:
if endpoint_resolver is None:
endpoint_resolver = resolve_public_push_endpoint
resolved = await endpoint_resolver(subscription["endpoint"])
if not resolved.addresses:
raise UnsafePushEndpoint("Endpoint must resolve to a public Web Push service")
if webpush_sender is None:
from pywebpush import webpush
webpush_sender = webpush
session = requests.Session()
session.trust_env = False
session.max_redirects = 0
origin = f"https://{resolved.hostname}"
session.mount(origin, _PinnedHTTPSAdapter(resolved))
await asyncio.to_thread(
webpush_sender,
subscription_info=subscription,
data=payload,
vapid_private_key=configuration.private_key,
vapid_claims={"sub": configuration.subject},
ttl=300,
timeout=10,
requests_session=session,
)
async def dispatch_unread_updates(
store: PushSubscriptionStore,
configuration: PushConfiguration,
unread: Callable[[], Awaitable[dict]],
send: Callable[[dict, str], Awaitable[None]] | None = None,
*,
session_statuses: Callable[[list[str]], Awaitable[dict[str, str]]] | None = None,
lease_seconds: float = 60.0,
send_timeout_seconds: float = 10.0,
max_concurrency: int = 8,
max_individual_notifications: int = 3,
endpoint_validator: Callable[[str], Awaitable[str]] | None = None,
) -> int:
if not configuration.enabled:
return 0
owner = secrets.token_urlsafe(18)
acquired = await asyncio.to_thread(
store.acquire_dispatch_lease,
owner,
channel="unread",
now=time.time(),
lease_seconds=lease_seconds,
)
if not acquired:
return 0
try:
page = await unread()
if page.get("complete") is False:
return 0
thread_revisions = {
int(item["id"]): str(item.get("updated_at") or "")
for item in page.get("items", [])
if isinstance(item, dict) and str(item.get("id", "")).isdigit()
}
await asyncio.to_thread(store.reconcile_unread, thread_revisions)
deliveries = await asyncio.to_thread(store.claim_unseen, thread_revisions)
if session_statuses is not None:
try:
statuses = await session_statuses(
[item.session_id for item in deliveries]
)
except Exception:
# Authorization state is mandatory for delivery. Preserve subscriptions
# so a temporary registry failure can be retried safely.
return 0
for delivery in deliveries:
if statuses.get(delivery.session_id) != "active":
await asyncio.to_thread(store.delete_session, delivery.session_id)
deliveries = [
delivery
for delivery in deliveries
if statuses.get(delivery.session_id) == "active"
]
semaphore = asyncio.Semaphore(max(1, max_concurrency))
ownership_lost = asyncio.Event()
async def dispatch_device(delivery) -> int:
async with semaphore:
if ownership_lost.is_set():
return 0
validate_endpoint = endpoint_validator
if validate_endpoint is None and send is None:
validate_endpoint = validate_public_push_endpoint
try:
if validate_endpoint is not None:
await validate_endpoint(delivery.subscription["endpoint"])
except UnsafePushEndpoint:
await asyncio.to_thread(store.delete_session, delivery.session_id)
return 0
count = 0
digest_pending = set(delivery.digest_revisions)
new_revisions = tuple(
thread_revision
for thread_revision in delivery.thread_revisions
if thread_revision not in digest_pending
)
individual_revisions = new_revisions[
:max(0, max_individual_notifications)
]
overflow_revisions = (
delivery.digest_revisions
+ new_revisions[len(individual_revisions):]
)
can_send_digest = True
for thread_id, revision in individual_revisions:
still_owner = await asyncio.to_thread(
store.acquire_dispatch_lease,
owner,
channel="unread",
now=time.time(),
lease_seconds=lease_seconds,
)
if not still_owner:
ownership_lost.set()
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}",
"notification_id": 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 isinstance(error, UnsafePushEndpoint) or status in {404, 410}:
await asyncio.to_thread(
store.delete_session, delivery.session_id
)
else:
await asyncio.to_thread(
store.mark_delivery_failed,
delivery.session_id,
"unread",
_delivery_failure_reason(error),
)
can_send_digest = False
# Leave this device's transient failures unseen for a later
# poll instead of paying the endpoint deadline repeatedly.
break
await asyncio.to_thread(
store.mark_delivery_succeeded, delivery.session_id, "unread"
)
await asyncio.to_thread(
store.mark_delivered,
delivery.session_id,
((thread_id, revision),),
)
count += 1
if overflow_revisions and can_send_digest:
still_owner = await asyncio.to_thread(
store.acquire_dispatch_lease,
owner,
channel="unread",
now=time.time(),
lease_seconds=lease_seconds,
)
if not still_owner:
ownership_lost.set()
return count
payload = json.dumps(
{
"title": f"{len(overflow_revisions)} new work updates",
"body": "Tap to review them in Stackchain.",
"route": "#/my-work/updates",
"tag": "stackchain-update-digest",
"update_count": len(overflow_revisions),
},
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 isinstance(error, UnsafePushEndpoint) or status in {404, 410}:
await asyncio.to_thread(
store.delete_session, delivery.session_id
)
else:
await asyncio.to_thread(
store.mark_delivery_failed,
delivery.session_id,
"unread",
_delivery_failure_reason(error),
)
await asyncio.to_thread(
store.mark_digest_pending,
delivery.session_id,
overflow_revisions,
)
return count
await asyncio.to_thread(
store.mark_delivery_succeeded, delivery.session_id, "unread"
)
await asyncio.to_thread(
store.mark_delivered,
delivery.session_id,
overflow_revisions,
)
count += 1
return count
async def dispatch_device_safely(delivery) -> int:
try:
return await dispatch_device(delivery)
except Exception as error:
# Device-specific validation, persistence, or provider failures
# must not cancel healthy siblings. Leave any uncheckpointed
# revisions unseen so a later poll can retry them.
try:
await asyncio.to_thread(
store.mark_delivery_failed,
delivery.session_id,
"unread",
_delivery_failure_reason(error),
)
except Exception:
pass
return 0
counts = await asyncio.gather(
*(dispatch_device_safely(delivery) for delivery in deliveries)
)
return sum(counts)
finally:
await asyncio.to_thread(store.release_dispatch_lease, owner, channel="unread")
async def dispatch_deadline_reminders(
store: PushSubscriptionStore,
configuration: PushConfiguration,
assigned: Callable[[], Awaitable[dict]],
send: Callable[[dict, str], Awaitable[None]] | None = None,
**kwargs,
) -> int:
owner = secrets.token_urlsafe(18)
acquired = await asyncio.to_thread(
store.acquire_dispatch_lease,
owner,
channel="deadline",
now=time.time(),
lease_seconds=max(
15.0,
float(kwargs.get("lease_seconds", 60.0)),
float(kwargs.get("send_timeout_seconds", 10.0)) + 5.0,
),
)
if not acquired:
return 0
try:
return await _dispatch_deadline_reminders_unlocked(
store, configuration, assigned, send, owner=owner, **kwargs
)
finally:
await asyncio.to_thread(store.release_dispatch_lease, owner, channel="deadline")
async def _dispatch_deadline_reminders_unlocked(
store: PushSubscriptionStore,
configuration: PushConfiguration,
assigned: Callable[[], Awaitable[dict]],
send: Callable[[dict, str], Awaitable[None]] | None = None,
*,
now: datetime | None = None,
session_statuses: Callable[[list[str]], Awaitable[dict[str, str]]] | None = None,
send_timeout_seconds: float = 10.0,
lease_seconds: float = 60.0,
max_concurrency: int = 8,
owner: str,
) -> int:
"""Send one privacy-safe Agenda digest per eligible device and local day."""
if not configuration.enabled:
return 0
devices = await asyncio.to_thread(store.deadline_reminder_devices)
if not devices:
return 0
current = now or datetime.now(timezone.utc)
eligible_devices = []
for device in devices:
try:
local_now = current.astimezone(ZoneInfo(device.timezone))
except ZoneInfoNotFoundError:
continue
snooze_due = (
device.snoozed_until is not None
and device.snoozed_until <= current.timestamp()
)
daily_due = (
device.snoozed_until is None
and local_now.hour >= device.reminder_hour
and device.delivered_local_day != local_now.date().isoformat()
)
if snooze_due or daily_due:
eligible_devices.append(device)
if not eligible_devices:
return 0
snapshot = await assigned()
if snapshot.get("complete") is False:
return 0
due_days = []
for item in snapshot.get("items", []):
if not isinstance(item, dict) or not item.get("due_date"):
continue
raw_due = str(item["due_date"])
try:
due_day = datetime.strptime(raw_due[:10], "%Y-%m-%d").date()
except ValueError:
continue
due_days.append(due_day)
if not due_days:
await asyncio.gather(*(
asyncio.to_thread(store.clear_deadline_snooze, device.session_id)
for device in eligible_devices
if device.snoozed_until is not None
))
return 0
due_counts = {}
for device in eligible_devices:
local_now = current.astimezone(ZoneInfo(device.timezone))
local_cutoff = local_now.date() + timedelta(days=device.reminder_days)
due_count = sum(due_day <= local_cutoff for due_day in due_days)
if due_count:
due_counts[device.session_id] = due_count
await asyncio.gather(*(
asyncio.to_thread(store.clear_deadline_snooze, device.session_id)
for device in eligible_devices
if device.snoozed_until is not None and device.session_id not in due_counts
))
eligible_devices = [
device for device in eligible_devices if device.session_id in due_counts
]
if not eligible_devices:
return 0
if session_statuses is not None:
try:
statuses = await session_statuses(
[device.session_id for device in eligible_devices]
)
except Exception:
return 0
for device in eligible_devices:
if statuses.get(device.session_id) != "active":
await asyncio.to_thread(store.delete_session, device.session_id)
eligible_devices = [
device
for device in eligible_devices
if statuses.get(device.session_id) == "active"
]
if not eligible_devices:
return 0
semaphore = asyncio.Semaphore(max(1, max_concurrency))
async def dispatch_device(device) -> int:
async with semaphore:
try:
local_now = current.astimezone(ZoneInfo(device.timezone))
except ZoneInfoNotFoundError:
return 0
local_day = local_now.date().isoformat()
snooze_due = (
device.snoozed_until is not None
and device.snoozed_until <= current.timestamp()
)
daily_due = (
device.snoozed_until is None
and local_now.hour >= device.reminder_hour
and device.delivered_local_day != local_day
)
if not (snooze_due or daily_due):
return 0
due_count = due_counts[device.session_id]
still_owner = await asyncio.to_thread(
store.acquire_dispatch_lease,
owner,
channel="deadline",
now=time.time(),
lease_seconds=max(15.0, lease_seconds, send_timeout_seconds + 5.0),
)
if not still_owner:
return 0
payload = json.dumps({
"title": f"{due_count} deadline{'s' if due_count != 1 else ''} need{'s' if due_count == 1 else ''} attention",
"body": f"Open Agenda to review or replan {'it' if due_count == 1 else 'them'}.",
"route": "#/my-work/agenda",
"protect_route": "#/my-work/agenda/protect-today",
"tag": f"stackchain-deadline-digest-{local_day}",
"deadline_count": due_count,
}, separators=(",", ":"))
try:
operation = (
send(device.subscription, payload)
if send is not None
else send_web_push(device.subscription, payload, configuration)
)
await asyncio.wait_for(operation, timeout=send_timeout_seconds)
except Exception as error:
status = getattr(getattr(error, "response", None), "status_code", None)
if isinstance(error, UnsafePushEndpoint) or status in {404, 410}:
await asyncio.to_thread(store.delete_session, device.session_id)
else:
await asyncio.to_thread(
store.mark_delivery_failed,
device.session_id,
"deadline",
_delivery_failure_reason(error),
)
return 0
await asyncio.to_thread(
store.mark_delivery_succeeded, device.session_id, "deadline"
)
await asyncio.to_thread(
store.mark_deadline_reminder_delivered, device.session_id, local_day
)
return 1
results = await asyncio.gather(
*(dispatch_device(device) for device in eligible_devices), return_exceptions=True
)
return sum(result for result in results if isinstance(result, int))