# backend/push.py # Push Notifications — Web Push API with VAPID authentication # ROADMAP #67 import base64 import json import logging import os from datetime import datetime, timezone from pathlib import Path from typing import Any from fastapi import APIRouter, Depends, HTTPException from pydantic import BaseModel, Field from backend.auth.middleware import require_auth logger = logging.getLogger("obsigate.push") router = APIRouter(prefix="/api/push", tags=["push"]) # VAPID keys storage VAPID_KEYS_FILE = Path(os.environ.get("OBSIGATE_DATA_DIR", "data")) / "vapid_keys.json" PUSH_SUBSCRIPTIONS_FILE = Path(os.environ.get("OBSIGATE_DATA_DIR", "data")) / "push_subscriptions.json" # In-memory cache _vapid_keys: dict[str, str] = {} _push_subscriptions: list[dict[str, Any]] = [] def load_vapid_keys() -> dict[str, str]: """Load VAPID keys from file or generate new ones.""" global _vapid_keys if _vapid_keys: return _vapid_keys try: if VAPID_KEYS_FILE.exists(): with open(VAPID_KEYS_FILE, "r") as f: _vapid_keys = json.load(f) else: # Generate new VAPID keys from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import ec private_key = ec.generate_private_key(ec.SECP256R1()) public_key = private_key.public_key() private_pem = private_key.private_bytes( encoding=serialization.Encoding.PEM, format=serialization.PrivateFormat.PKCS8, encryption_algorithm=serialization.NoEncryption() ).decode('utf-8') public_pem = public_key.public_bytes( encoding=serialization.Encoding.PEM, format=serialization.PublicFormat.SubjectPublicKeyInfo ).decode('utf-8') # Convert to base64url for Web Push public_numbers = public_key.public_numbers() def int_to_base64url(n: int) -> str: byte_len = (n.bit_length() + 7) // 8 return base64.urlsafe_b64encode(n.to_bytes(byte_len, 'big')).decode('utf-8').rstrip('=') _vapid_keys = { "private_key": private_pem, "public_key": public_pem, "public_key_base64url": int_to_base64url(public_numbers.x) + "." + int_to_base64url(public_numbers.y) } VAPID_KEYS_FILE.parent.mkdir(parents=True, exist_ok=True) with open(VAPID_KEYS_FILE, "w") as f: json.dump(_vapid_keys, f) except Exception as e: logger.error(f"Failed to load/generate VAPID keys: {e}") # Fallback: return empty (will use demo keys if needed) _vapid_keys = {} return _vapid_keys def get_vapid_public_key() -> str: """Get VAPID public key in base64url format for client.""" keys = load_vapid_keys() return keys.get("public_key_base64url", "") def load_push_subscriptions() -> list[dict[str, Any]]: """Load push subscriptions from file.""" global _push_subscriptions if _push_subscriptions: return _push_subscriptions try: if PUSH_SUBSCRIPTIONS_FILE.exists(): with open(PUSH_SUBSCRIPTIONS_FILE, "r") as f: _push_subscriptions = json.load(f) except Exception as e: logger.error(f"Failed to load push subscriptions: {e}") _push_subscriptions = [] return _push_subscriptions def save_push_subscriptions() -> None: """Save push subscriptions to file.""" try: PUSH_SUBSCRIPTIONS_FILE.parent.mkdir(parents=True, exist_ok=True) with open(PUSH_SUBSCRIPTIONS_FILE, "w") as f: json.dump(_push_subscriptions, f, indent=2) except Exception as e: logger.error(f"Failed to save push subscriptions: {e}") # ── Models ────────────────────────────────────────────────────────────────── class PushSubscription(BaseModel): """Push subscription from client (matches Push API).""" endpoint: str keys: dict[str, str] # { p256dh, auth } class SubscribeRequest(BaseModel): """Request to subscribe to push notifications.""" subscription: PushSubscription vault: str = Field(description="Vault name this subscription is for") class SubscribeResponse(BaseModel): """Response to subscription request.""" success: bool subscription_id: str | None = None class VapidPublicKeyResponse(BaseModel): """VAPID public key for client.""" public_key: str class PushPayload(BaseModel): """Payload for sending a push notification.""" vault: str title: str body: str data: dict[str, Any] = Field(default_factory=dict) tag: str = "obsigate-notification" # ── Endpoints ─────────────────────────────────────────────────────────────── @router.get("/vapid-public-key", response_model=VapidPublicKeyResponse) async def get_vapid_public_key_endpoint(): """Get VAPID public key for client subscription.""" public_key = get_vapid_public_key() if not public_key: # Return a demo key if generation failed (for testing) return {"public_key": "demo-key-for-testing"} return {"public_key": public_key} @router.post("/subscribe", response_model=SubscribeResponse) async def subscribe_push( request: SubscribeRequest, current_user=Depends(require_auth) ): """Subscribe to push notifications for a vault.""" username = current_user.get("username", "unknown") # Check if user has access to this vault user_vaults = current_user.get("_token_vaults") or current_user.get("vaults", []) if "*" not in user_vaults and request.vault not in user_vaults: raise HTTPException(status_code=403, detail="No access to this vault") # Check if subscription already exists for sub in _push_subscriptions: if sub["endpoint"] == request.subscription.endpoint and sub["username"] == username: return SubscribeResponse(success=True, subscription_id=sub.get("id")) # Add new subscription sub_id = base64.urlsafe_b64encode(os.urandom(16)).decode('utf-8').rstrip('=') subscription = { "id": sub_id, "endpoint": request.subscription.endpoint, "keys": request.subscription.keys, "vault": request.vault, "username": username, "created_at": datetime.now(timezone.utc).isoformat() } _push_subscriptions.append(subscription) save_push_subscriptions() logger.info(f"Push subscription added for user={username}, vault={request.vault}") return SubscribeResponse(success=True, subscription_id=sub_id) @router.delete("/subscribe") async def unsubscribe_push( endpoint: str, current_user=Depends(require_auth) ): """Unsubscribe from push notifications.""" username = current_user.get("username", "unknown") global _push_subscriptions original_len = len(_push_subscriptions) _push_subscriptions = [ s for s in _push_subscriptions if not (s["endpoint"] == endpoint and s["username"] == username) ] if len(_push_subscriptions) < original_len: save_push_subscriptions() return {"success": True, "message": "Unsubscribed"} return {"success": False, "message": "Subscription not found"} @router.get("/subscriptions") async def list_subscriptions(current_user=Depends(require_auth)): """List current user's push subscriptions.""" username = current_user.get("username", "unknown") user_subs = [ { "id": s["id"], "vault": s["vault"], "created_at": s["created_at"], "endpoint": s["endpoint"][:50] + "..." # Truncate for privacy } for s in _push_subscriptions if s["username"] == username ] return {"subscriptions": user_subs} # ── Internal: Send push notification ──────────────────────────────────────── async def send_push_notification(vault: str, title: str, body: str, data: dict | None = None, tag: str = "obsigate-notification") -> int: """Send push notification to all subscribers of a vault. Returns count of sent notifications.""" subscriptions = [s for s in _push_subscriptions if s["vault"] == vault] if not subscriptions: return 0 payload = { "title": title, "body": body, "data": data or {}, "tag": tag } # Import here to avoid circular dependency from pywebpush import WebPushException, webpush keys = load_vapid_keys() vapid_private_key = keys.get("private_key") vapid_claims = { "sub": "mailto:admin@obsigate.local" } sent = 0 for sub in subscriptions: try: webpush( subscription_info={ "endpoint": sub["endpoint"], "keys": sub["keys"] }, data=json.dumps(payload), vapid_private_key=vapid_private_key, vapid_claims=vapid_claims ) sent += 1 except WebPushException as e: logger.warning(f"Push failed for subscription {sub['id']}: {e}") # If subscription expired/invalid, remove it if e.response and e.response.status_code in (404, 410): _push_subscriptions.remove(sub) except Exception as e: logger.error(f"Push error for {sub['id']}: {e}") if sent < len(subscriptions): save_push_subscriptions() return sent