aboutsummaryrefslogtreecommitdiffstats
path: root/packages/meshbay-hub/src/meshbay_hub/webpush.py
diff options
context:
space:
mode:
Diffstat (limited to 'packages/meshbay-hub/src/meshbay_hub/webpush.py')
-rw-r--r--packages/meshbay-hub/src/meshbay_hub/webpush.py194
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