# retoor import base64 import fcntl import json import logging import os import random import time from datetime import datetime, timezone from typing import Any from urllib.parse import urlparse import httpx import jwt from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric import ec from cryptography.hazmat.primitives.ciphers.aead import AESGCM from cryptography.hazmat.primitives.hashes import SHA256 from cryptography.hazmat.primitives.kdf.hkdf import HKDF from devplacepy import stealth from devplacepy.config import ( SECONDS_PER_DAY, VAPID_PRIVATE_KEY_FILE, VAPID_PRIVATE_KEY_PKCS8_FILE, VAPID_PUBLIC_KEY_FILE, VAPID_SUB, ) from devplacepy.database import get_table from devplacepy.utils import generate_uid logger = logging.getLogger(__name__) JWT_LIFETIME_SECONDS = 60 * 60 PUSH_TTL_SECONDS = str(SECONDS_PER_DAY) DEAD_SUBSCRIPTION_STATUSES = (404, 410) ACCEPTED_STATUSES = (200, 201) def generate_private_key() -> None: if not VAPID_PRIVATE_KEY_FILE.exists(): private_key = ec.generate_private_key(ec.SECP256R1(), default_backend()) pem = private_key.private_bytes( encoding=serialization.Encoding.PEM, format=serialization.PrivateFormat.TraditionalOpenSSL, encryption_algorithm=serialization.NoEncryption(), ) VAPID_PRIVATE_KEY_FILE.write_bytes(pem) logger.info("Generated VAPID private key at %s", VAPID_PRIVATE_KEY_FILE) def generate_pkcs8_private_key() -> None: if not VAPID_PRIVATE_KEY_PKCS8_FILE.exists(): private_key = serialization.load_pem_private_key( VAPID_PRIVATE_KEY_FILE.read_bytes(), password=None, backend=default_backend(), ) pem = private_key.private_bytes( encoding=serialization.Encoding.PEM, format=serialization.PrivateFormat.PKCS8, encryption_algorithm=serialization.NoEncryption(), ) VAPID_PRIVATE_KEY_PKCS8_FILE.write_bytes(pem) logger.info( "Generated VAPID PKCS8 private key at %s", VAPID_PRIVATE_KEY_PKCS8_FILE ) def generate_public_key() -> None: if not VAPID_PUBLIC_KEY_FILE.exists(): private_key = serialization.load_pem_private_key( VAPID_PRIVATE_KEY_FILE.read_bytes(), password=None, backend=default_backend(), ) pem = private_key.public_key().public_bytes( encoding=serialization.Encoding.PEM, format=serialization.PublicFormat.SubjectPublicKeyInfo, ) VAPID_PUBLIC_KEY_FILE.write_bytes(pem) logger.info("Generated VAPID public key at %s", VAPID_PUBLIC_KEY_FILE) def ensure_certificates() -> None: if ( VAPID_PRIVATE_KEY_FILE.exists() and VAPID_PRIVATE_KEY_PKCS8_FILE.exists() and VAPID_PUBLIC_KEY_FILE.exists() ): return lock_path = VAPID_PRIVATE_KEY_FILE.parent / ".vapid.lock" with open(lock_path, "w") as lock: fcntl.flock(lock, fcntl.LOCK_EX) try: generate_private_key() generate_pkcs8_private_key() generate_public_key() finally: fcntl.flock(lock, fcntl.LOCK_UN) def hkdf(input_key: bytes, salt: bytes, info: bytes, length: int) -> bytes: return HKDF( algorithm=SHA256(), length=length, salt=salt, info=info, backend=default_backend(), ).derive(input_key) def browser_base64(data: bytes) -> str: return base64.urlsafe_b64encode(data).decode("utf-8").rstrip("=") _keys: dict[str, Any] = {} def _load_keys() -> dict[str, Any]: if _keys: return _keys ensure_certificates() private_key = serialization.load_pem_private_key( VAPID_PRIVATE_KEY_FILE.read_bytes(), password=None, backend=default_backend() ) public_key = serialization.load_pem_public_key( VAPID_PUBLIC_KEY_FILE.read_bytes(), backend=default_backend() ) uncompressed_point = public_key.public_bytes( encoding=serialization.Encoding.X962, format=serialization.PublicFormat.UncompressedPoint, ) _keys["private_key_pem"] = private_key.private_bytes( encoding=serialization.Encoding.PEM, format=serialization.PrivateFormat.TraditionalOpenSSL, encryption_algorithm=serialization.NoEncryption(), ) _keys["public_key_point"] = uncompressed_point _keys["public_key_base64"] = browser_base64(uncompressed_point) logger.debug("Loaded VAPID key material into cache") return _keys def public_key_standard_b64() -> str: point = _load_keys()["public_key_point"] return base64.b64encode(point).decode("utf-8").rstrip("=") def create_notification_authorization(push_url: str) -> str: target = urlparse(push_url) audience = f"{target.scheme}://{target.netloc}" issued_at = int(time.time()) return jwt.encode( { "sub": VAPID_SUB, "aud": audience, "exp": issued_at + JWT_LIFETIME_SECONDS, "nbf": issued_at, "iat": issued_at, "jti": generate_uid(), }, _load_keys()["private_key_pem"], algorithm="ES256", ) def create_notification_info_with_payload( endpoint: str, auth: str, p256dh: str, payload: str ) -> dict[str, Any]: message_private_key = ec.generate_private_key(ec.SECP256R1(), default_backend()) message_public_key_bytes = message_private_key.public_key().public_bytes( encoding=serialization.Encoding.X962, format=serialization.PublicFormat.UncompressedPoint, ) salt = os.urandom(16) user_key_bytes = base64.urlsafe_b64decode(p256dh + "==") shared_secret = message_private_key.exchange( ec.ECDH(), ec.EllipticCurvePublicKey.from_encoded_point(ec.SECP256R1(), user_key_bytes), ) encryption_key = hkdf( shared_secret, base64.urlsafe_b64decode(auth + "=="), b"Content-Encoding: auth\x00", 32, ) context = ( b"P-256\x00" + len(user_key_bytes).to_bytes(2, "big") + user_key_bytes + len(message_public_key_bytes).to_bytes(2, "big") + message_public_key_bytes ) nonce = hkdf(encryption_key, salt, b"Content-Encoding: nonce\x00" + context, 12) content_encryption_key = hkdf( encryption_key, salt, b"Content-Encoding: aesgcm\x00" + context, 16 ) padding_length = random.randint(0, 16) padding = padding_length.to_bytes(2, "big") + b"\x00" * padding_length data = AESGCM(content_encryption_key).encrypt( nonce, padding + payload.encode("utf-8"), None ) return { "headers": { "Authorization": f"WebPush {create_notification_authorization(endpoint)}", "Crypto-Key": f"dh={browser_base64(message_public_key_bytes)}; p256ecdsa={_load_keys()['public_key_base64']}", "Encryption": f"salt={browser_base64(salt)}", "Content-Encoding": "aesgcm", "Content-Length": str(len(data)), "Content-Type": "application/octet-stream", }, "data": data, } def _mark_subscription_dead(subscription_id: int) -> None: get_table("push_registration").update( {"id": subscription_id, "deleted_at": datetime.now(timezone.utc).isoformat()}, ["id"], ) logger.info("Soft-deleted dead push subscription id=%s", subscription_id) async def notify_user(user_uid: str, payload: dict[str, Any]) -> None: registrations = list( get_table("push_registration").find(user_uid=user_uid, deleted_at=None) ) if not registrations: logger.debug("No active push subscriptions for user %s", user_uid) return body = json.dumps(payload) async with stealth.stealth_async_client(timeout=10.0) as client: for subscription in registrations: endpoint = subscription["endpoint"] try: notification_payload = create_notification_info_with_payload( endpoint, subscription["key_auth"], subscription["key_p256dh"], body, ) headers = {**notification_payload["headers"], "TTL": PUSH_TTL_SECONDS} response = await client.post( endpoint, headers=headers, content=notification_payload["data"] ) except (httpx.HTTPError, ValueError) as exc: logger.warning("Push error for %s via %s: %s", user_uid, endpoint, exc) continue if response.status_code in ACCEPTED_STATUSES: logger.debug("Push delivered to %s via %s", user_uid, endpoint) elif response.status_code in DEAD_SUBSCRIPTION_STATUSES: _mark_subscription_dead(subscription["id"]) else: logger.warning( "Push rejected (%s) for %s via %s", response.status_code, user_uid, endpoint, ) async def register( user_uid: str, endpoint: str, key_auth: str, key_p256dh: str ) -> tuple[dict[str, Any], bool]: table = get_table("push_registration") existing = table.find_one( user_uid=user_uid, endpoint=endpoint, key_auth=key_auth, key_p256dh=key_p256dh, deleted_at=None, ) if existing: logger.debug("Push subscription already registered for user %s", user_uid) return existing, False record = { "uid": generate_uid(), "user_uid": user_uid, "endpoint": endpoint, "key_auth": key_auth, "key_p256dh": key_p256dh, "created_at": datetime.now(timezone.utc).isoformat(), "deleted_at": None, } table.insert(record) logger.info("Registered push subscription for user %s", user_uid) return record, True