diff options
Diffstat (limited to 'packages/meshbay-hub/src/meshbay_hub/webpush.py')
| -rw-r--r-- | packages/meshbay-hub/src/meshbay_hub/webpush.py | 194 |
1 files changed, 194 insertions, 0 deletions
diff --git a/packages/meshbay-hub/src/meshbay_hub/webpush.py b/packages/meshbay-hub/src/meshbay_hub/webpush.py new file mode 100644 index 0000000..301e191 --- /dev/null +++ b/packages/meshbay-hub/src/meshbay_hub/webpush.py @@ -0,0 +1,194 @@ +""" +Web Push to a phone: RFC 8291 encryption and the one outbound request. + +The Android application registers with a UnifiedPush distributor (ntfy, or any +other) and hands the hub an endpoint URL and a P-256 key; the hub POSTs each +notification there, encrypted to that key (`docs/MESHBAY_DESIGN.md` §7.6). The +push server relays bytes it cannot read; what it does learn is *when* this +person is notified, which is the same metadata the hub already holds. + +**The endpoint is a URL a member supplied, and the hub fetches it.** That is +the shape of an SSRF, so a send resolves the host itself, refuses unless every +address is public, and connects to the address it checked — the hostname rides +only as the TLS server name and the Host header, so a second resolution cannot +point the request somewhere else. No redirect is followed. +""" + +import asyncio +import base64 +import ipaddress +import json +import logging +import os +import socket +from dataclasses import dataclass +from urllib.parse import urlsplit, urlunsplit + +import httpx +from cryptography.hazmat.primitives import hashes, hmac, serialization +from cryptography.hazmat.primitives.asymmetric import ec +from cryptography.hazmat.primitives.ciphers.aead import AESGCM + +log = logging.getLogger(__name__) + +RECORD_SIZE = 4096 +SEND_TIMEOUT = 5.0 +MAX_ENDPOINT = 1024 +# RFC 8030 §5.2: a push service may keep a message this long for a phone that is +# off. A chat line an hour old is still worth seeing; one a day old is not. +TTL_CHAT = 3600 +TTL_OTHER = 86400 + + +class EndpointRefused(ValueError): + """The endpoint is not one the hub will send to.""" + + +def b64url_decode(value: str) -> bytes: + return base64.urlsafe_b64decode(value + "=" * (-len(value) % 4)) + + +def _hmac(key: bytes, data: bytes) -> bytes: + h = hmac.HMAC(key, hashes.SHA256()) + h.update(data) + return h.finalize() + + +def check_keys(p256dh: str, auth: str) -> tuple[bytes, bytes]: + """Decode and validate a subscription's keys; ValueError when they are not.""" + try: + ua_public = b64url_decode(p256dh) + auth_secret = b64url_decode(auth) + except (ValueError, TypeError) as e: + raise ValueError("keys are not base64url") from e + if len(ua_public) != 65 or ua_public[0] != 4: + raise ValueError("p256dh is not an uncompressed P-256 point") + if len(auth_secret) != 16: + raise ValueError("auth is not 16 bytes") + # Raises ValueError for a point that is not on the curve. + ec.EllipticCurvePublicKey.from_encoded_point(ec.SECP256R1(), ua_public) + return ua_public, auth_secret + + +def encrypt(plaintext: bytes, ua_public: bytes, auth_secret: bytes, *, + as_private: ec.EllipticCurvePrivateKey | None = None, + salt: bytes | None = None) -> bytes: + """One aes128gcm record (RFC 8188) keyed as RFC 8291 §3.4 says. + + `as_private` and `salt` are parameters only so the RFC's own example can be + replayed; a send always draws both fresh. + """ + if len(plaintext) > RECORD_SIZE - 16 - 1 - 86: + raise ValueError("push payload too large for one record") + as_private = as_private or ec.generate_private_key(ec.SECP256R1()) + salt = salt or os.urandom(16) + as_public = as_private.public_key().public_bytes( + serialization.Encoding.X962, serialization.PublicFormat.UncompressedPoint) + ua_key = ec.EllipticCurvePublicKey.from_encoded_point(ec.SECP256R1(), ua_public) + ecdh_secret = as_private.exchange(ec.ECDH(), ua_key) + + prk_key = _hmac(auth_secret, ecdh_secret) + key_info = b"WebPush: info\x00" + ua_public + as_public + ikm = _hmac(prk_key, key_info + b"\x01") + prk = _hmac(salt, ikm) + cek = _hmac(prk, b"Content-Encoding: aes128gcm\x00\x01")[:16] + nonce = _hmac(prk, b"Content-Encoding: nonce\x00\x01")[:12] + + header = salt + RECORD_SIZE.to_bytes(4, "big") + bytes([len(as_public)]) + as_public + return header + AESGCM(cek).encrypt(nonce, plaintext + b"\x02", None) + + +def check_endpoint(url: str) -> tuple[str, int]: + """The endpoint's shape, checked when it is registered: https, a host, no + credentials, not an address that is private on its face. Where its name + resolves is checked at every send, because that can change.""" + if len(url) > MAX_ENDPOINT: + raise EndpointRefused("endpoint too long") + parts = urlsplit(url) + if parts.scheme != "https" or not parts.hostname: + raise EndpointRefused("endpoint must be an https URL") + if parts.username or parts.password: + raise EndpointRefused("endpoint must not carry credentials") + try: + port = parts.port or 443 + except ValueError as e: + raise EndpointRefused("endpoint port is not a number") from e + host = parts.hostname + try: + literal = ipaddress.ip_address(host) + except ValueError: + literal = None + if literal is not None and not literal.is_global: + raise EndpointRefused("endpoint is not a public address") + return host, port + + +async def _resolve_public(host: str, port: int) -> str: + infos = await asyncio.get_running_loop().getaddrinfo( + host, port, type=socket.SOCK_STREAM) + addresses = {info[4][0] for info in infos} + if not addresses: + raise EndpointRefused("endpoint does not resolve") + for a in addresses: + if not ipaddress.ip_address(a.split("%", 1)[0]).is_global: + raise EndpointRefused("endpoint resolves to a non-public address") + # IPv4 first: a hub with an AAAA answer and no IPv6 route is common. + return sorted(addresses, key=lambda a: (":" in a, a))[0] + + +@dataclass +class Target: + endpoint: str + p256dh: str + auth: str + + +# Outcomes of one send, for the caller to act on. +DELIVERED = "delivered" +GONE = "gone" # 404/410: the registration no longer exists (RFC 8030 §7.3) +FAILED = "failed" + + +async def send(target: Target, payload: dict, *, ttl: int, + client: httpx.AsyncClient | None = None) -> str: + """Encrypt `payload` for one subscription and POST it. Never raises.""" + try: + host, port = check_endpoint(target.endpoint) + ua_public, auth_secret = check_keys(target.p256dh, target.auth) + body = encrypt(json.dumps(payload, separators=(",", ":")).encode(), + ua_public, auth_secret) + address = await _resolve_public(host, port) + except (EndpointRefused, ValueError, OSError) as e: + log.info("push refused before sending: %s", e) + return FAILED + + parts = urlsplit(target.endpoint) + netloc = f"[{address}]" if ":" in address else address + if parts.port: + netloc += f":{parts.port}" + pinned = urlunsplit((parts.scheme, netloc, parts.path or "/", parts.query, "")) + headers = { + "Host": parts.netloc, + "Content-Encoding": "aes128gcm", + "Content-Type": "application/octet-stream", + "TTL": str(ttl), + "Urgency": "normal", + } + own = client is None + client = client or httpx.AsyncClient(timeout=SEND_TIMEOUT, follow_redirects=False) + try: + resp = await client.post(pinned, content=body, headers=headers, + extensions={"sni_hostname": host}) + except httpx.HTTPError as e: + # The URL path is a bearer capability for this phone: never logged. + log.info("push to %s failed: %s", host, type(e).__name__) + return FAILED + finally: + if own: + await client.aclose() + if resp.status_code in (404, 410): + return GONE + if 200 <= resp.status_code < 300: + return DELIVERED + log.info("push to %s answered %d", host, resp.status_code) + return FAILED |