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))