#!/usr/bin/env -S uv run --quiet --script # /// script # requires-python = ">=3.11" # dependencies = ["noiseprotocol>=0.3.1", "cryptography>=42"] # /// """Smol Mail reference server. Implements SPEC.md version 1: a store-and-forward mailbox reachable over TCP with a Noise_NX handshake. The server never sees plaintext, sender identities, or any private key belonging to a user; it stores sealed envelopes addressed to a recipient public key and hands them back to whoever can sign for that key. uv run smolmaild.py keygen --key server.key uv run smolmaild.py serve --key server.key --db mail.db """ from __future__ import annotations import argparse import base64 import hashlib import logging import os import socket import socketserver import sqlite3 import struct import sys import threading import time from cryptography.exceptions import InvalidSignature from cryptography.hazmat.primitives import serialization from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PublicKey from cryptography.hazmat.primitives.asymmetric.x25519 import X25519PrivateKey from noise.connection import Keypair, NoiseConnection log = logging.getLogger("smolmaild") # --------------------------------------------------------------------------- # Protocol constants (SPEC.md §4, §6, §11, §12) # --------------------------------------------------------------------------- NOISE_PROTOCOL = b"Noise_NX_25519_ChaChaPoly_SHA256" PROLOGUE = b"smolmail/1" LABEL_AUTH = b"smolmail/1 auth" LABEL_ID = b"smolmail/1 id" LABEL_ROTATE = b"smolmail/1 rotate" OP_AUTH = 0x00 OP_RESOLVE = 0x01 OP_SEND = 0x02 OP_FETCH = 0x03 OP_DELETE = 0x04 OP_REGISTER = 0x05 OK = 0 MALFORMED = 1 BAD_VERSION = 2 UNKNOWN_USER = 3 AUTH_REQUIRED = 4 AUTH_FAILED = 5 QUOTA_EXCEEDED = 6 TOO_LARGE = 7 RATE_LIMITED = 8 NOT_PERMITTED = 9 INTERNAL_ERROR = 10 ENVELOPE_MAGIC = b"SMOL" ENVELOPE_VERSION = 1 ENVELOPE_HEADER = 69 # magic 4 + version 1 + to 32 + epk 32 ENVELOPE_MIN = ENVELOPE_HEADER + 16 # + Poly1305 tag ID_LEN = 16 KEY_LEN = 32 CERT_LEN = 136 # old_pub 32 + new_pub 32 + time 8 + signature 64 MAX_CHAIN = 16 DEFAULT_PORT = 1961 MAX_FRAME = 1 << 20 # §4 application frame ceiling NOISE_MAX = 65535 # §4 Noise message ceiling NOISE_PAYLOAD = NOISE_MAX - 16 # minus the AEAD tag FETCH_BUDGET = 512 * 1024 # §6.1, must stay under MAX_FRAME IDLE_TIMEOUT = 120.0 PURGE_INTERVAL = 60.0 class ProtocolError(Exception): """A peer sent something unparseable. Always answered with MALFORMED.""" # --------------------------------------------------------------------------- # Encoding helpers # --------------------------------------------------------------------------- def b32(raw: bytes) -> str: """RFC 4648 base32, lowercase and unpadded, as §3 requires.""" return base64.b32encode(raw).decode("ascii").rstrip("=").lower() class Reader: """Fail-closed reader over a frame body. Every parse path raises rather than reading past the end, so a truncated frame can never be mistaken for a short but valid one. """ def __init__(self, buf: bytes) -> None: self.buf = buf self.pos = 0 def take(self, n: int) -> bytes: if n < 0 or self.pos + n > len(self.buf): raise ProtocolError(f"short read: want {n}, have {len(self.buf) - self.pos}") out = self.buf[self.pos : self.pos + n] self.pos += n return out def u8(self) -> int: return self.take(1)[0] def u16(self) -> int: return struct.unpack(">H", self.take(2))[0] def u32(self) -> int: return struct.unpack(">I", self.take(4))[0] def i64(self) -> int: return struct.unpack(">q", self.take(8))[0] def rest(self) -> bytes: return self.take(len(self.buf) - self.pos) def done(self) -> None: if self.pos != len(self.buf): raise ProtocolError(f"{len(self.buf) - self.pos} trailing bytes") def valid_username(name: str) -> bool: """§3: 1-63 bytes of [a-z0-9._-], not starting or ending with a separator.""" if not 1 <= len(name) <= 63: return False if not all(c.isdigit() or ("a" <= c <= "z") or c in "._-" for c in name): return False return name[0] not in "._-" and name[-1] not in "._-" def message_id(envelope: bytes) -> bytes: """§5.4. Derived from the envelope, so a sender cannot choose it.""" return hashlib.sha256(LABEL_ID + envelope).digest()[:ID_LEN] def verify(pubkey: bytes, signature: bytes, message: bytes) -> bool: try: Ed25519PublicKey.from_public_bytes(pubkey).verify(signature, message) return True except (InvalidSignature, ValueError): return False # --------------------------------------------------------------------------- # Storage # --------------------------------------------------------------------------- SCHEMA = """ CREATE TABLE IF NOT EXISTS users ( username TEXT PRIMARY KEY, identity BLOB NOT NULL ); -- Every key ever bound to a username, so a superseded key stays addressable -- across a rotation (§7). CREATE TABLE IF NOT EXISTS keys ( identity BLOB PRIMARY KEY, username TEXT NOT NULL ); CREATE TABLE IF NOT EXISTS rotations ( username TEXT NOT NULL, seq INTEGER NOT NULL, cert BLOB NOT NULL, PRIMARY KEY (username, seq) ); CREATE TABLE IF NOT EXISTS messages ( id BLOB PRIMARY KEY, recipient BLOB NOT NULL, received_at INTEGER NOT NULL, envelope BLOB NOT NULL ); CREATE INDEX IF NOT EXISTS messages_by_recipient ON messages (recipient, received_at); """ class Store: """SQLite behind a connection per thread. ThreadingTCPServer hands each connection its own thread and sqlite3 connections are not shareable across threads, so they are kept in thread-local state rather than guarded by a lock. """ def __init__(self, path: str) -> None: self.path = path self.local = threading.local() with sqlite3.connect(path) as db: db.execute("PRAGMA journal_mode=WAL") db.executescript(SCHEMA) @property def db(self) -> sqlite3.Connection: conn = getattr(self.local, "conn", None) if conn is None: conn = sqlite3.connect(self.path, timeout=10.0) conn.execute("PRAGMA busy_timeout=10000") self.local.conn = conn return conn def identity_of(self, username: str) -> bytes | None: row = self.db.execute( "SELECT identity FROM users WHERE username = ?", (username,) ).fetchone() return row[0] if row else None def chain(self, username: str) -> list[bytes]: return [ row[0] for row in self.db.execute( "SELECT cert FROM rotations WHERE username = ? ORDER BY seq", (username,) ) ] def keys_of(self, username: str) -> list[bytes]: return [ row[0] for row in self.db.execute( "SELECT identity FROM keys WHERE username = ?", (username,) ) ] def username_for_key(self, identity: bytes) -> str | None: row = self.db.execute( "SELECT username FROM keys WHERE identity = ?", (identity,) ).fetchone() return row[0] if row else None def register(self, username: str, identity: bytes) -> bool: try: with self.db: self.db.execute( "INSERT INTO users (username, identity) VALUES (?, ?)", (username, identity), ) self.db.execute( "INSERT INTO keys (identity, username) VALUES (?, ?)", (identity, username), ) return True except sqlite3.IntegrityError: # Username taken, or this key is already bound to another username. return False def rotate(self, username: str, new_key: bytes, cert: bytes, seq: int) -> bool: try: with self.db: self.db.execute( "UPDATE users SET identity = ? WHERE username = ?", (new_key, username), ) self.db.execute( "INSERT INTO keys (identity, username) VALUES (?, ?)", (new_key, username), ) self.db.execute( "INSERT INTO rotations (username, seq, cert) VALUES (?, ?, ?)", (username, seq, cert), ) return True except sqlite3.IntegrityError: return False def mailbox_bytes(self, keys: list[bytes]) -> int: marks = ",".join("?" * len(keys)) row = self.db.execute( f"SELECT COALESCE(SUM(LENGTH(envelope)), 0) FROM messages " f"WHERE recipient IN ({marks})", keys, ).fetchone() return row[0] def store_message(self, mid: bytes, recipient: bytes, envelope: bytes) -> None: with self.db: self.db.execute( "INSERT OR IGNORE INTO messages (id, recipient, received_at, envelope) " "VALUES (?, ?, ?, ?)", (mid, recipient, int(time.time()), envelope), ) def pending(self, keys: list[bytes], budget: int) -> list[tuple[bytes, int, bytes]]: marks = ",".join("?" * len(keys)) rows = self.db.execute( f"SELECT id, received_at, envelope FROM messages " f"WHERE recipient IN ({marks}) ORDER BY received_at, id", keys, ) out: list[tuple[bytes, int, bytes]] = [] used = 0 for mid, received_at, envelope in rows: # Always return at least one message, even if it alone exceeds the # budget, so an oversized envelope cannot wedge a mailbox shut. if out and used + len(envelope) > budget: break out.append((mid, received_at, envelope)) used += len(envelope) return out def delete(self, keys: list[bytes], ids: list[bytes]) -> int: kmarks = ",".join("?" * len(keys)) imarks = ",".join("?" * len(ids)) with self.db: cur = self.db.execute( f"DELETE FROM messages WHERE id IN ({imarks}) " f"AND recipient IN ({kmarks})", ids + keys, ) return cur.rowcount def purge(self, older_than: int) -> int: with self.db: cur = self.db.execute("DELETE FROM messages WHERE received_at < ?", (older_than,)) return cur.rowcount # --------------------------------------------------------------------------- # Rate limiting # --------------------------------------------------------------------------- class RateLimiter: """Fixed-window per-IP counter, the whole of §10's abuse control. A server cannot see senders, so quotas, size caps and this are all it has. """ def __init__(self, limit: int, window: float = 60.0) -> None: self.limit = limit self.window = window self.lock = threading.Lock() self.hits: dict[str, tuple[float, int]] = {} def allow(self, ip: str) -> bool: if self.limit <= 0: return True now = time.monotonic() with self.lock: start, count = self.hits.get(ip, (now, 0)) if now - start >= self.window: start, count = now, 0 if count >= self.limit: return False self.hits[ip] = (start, count + 1) if len(self.hits) > 4096: self._evict(now) return True def _evict(self, now: float) -> None: stale = [ip for ip, (start, _) in self.hits.items() if now - start >= self.window] for ip in stale: del self.hits[ip] # --------------------------------------------------------------------------- # Framed transport (§4) # --------------------------------------------------------------------------- def recv_exact(sock: socket.socket, n: int) -> bytes: chunks = [] remaining = n while remaining: chunk = sock.recv(min(remaining, 65536)) if not chunk: raise EOFError("peer closed") chunks.append(chunk) remaining -= len(chunk) return b"".join(chunks) class Channel: """Noise messages under a u16 length prefix, application frames on top. The two layers are independent: an application frame is split across as many Noise messages as it needs, and reassembled from them. """ def __init__(self, sock: socket.socket, noise: NoiseConnection) -> None: self.sock = sock self.noise = noise self.buf = bytearray() def _read_noise(self) -> bytes: (length,) = struct.unpack(">H", recv_exact(self.sock, 2)) return self.noise.decrypt(recv_exact(self.sock, length)) def _write_noise(self, payload: bytes) -> None: packet = self.noise.encrypt(payload) self.sock.sendall(struct.pack(">H", len(packet)) + packet) def read_frame(self) -> tuple[int, bytes]: while len(self.buf) < 5: self.buf += self._read_noise() (length,) = struct.unpack(">I", self.buf[:4]) if length < 1 or length > MAX_FRAME: raise ProtocolError(f"frame length {length} out of range") while len(self.buf) < 4 + length: self.buf += self._read_noise() frame = bytes(self.buf[4 : 4 + length]) del self.buf[: 4 + length] return frame[0], frame[1:] def write_frame(self, op: int, body: bytes) -> None: frame = struct.pack(">I", 1 + len(body)) + bytes([op]) + body for off in range(0, len(frame), NOISE_PAYLOAD): self._write_noise(frame[off : off + NOISE_PAYLOAD]) def handshake(sock: socket.socket, static_key: bytes) -> NoiseConnection: """Noise_NX responder. The initiator stays anonymous; only we hold a static.""" noise = NoiseConnection.from_name(NOISE_PROTOCOL) noise.set_as_responder() noise.set_prologue(PROLOGUE) noise.set_keypair_from_private_bytes(Keypair.STATIC, static_key) noise.start_handshake() (length,) = struct.unpack(">H", recv_exact(sock, 2)) noise.read_message(recv_exact(sock, length)) reply = noise.write_message() sock.sendall(struct.pack(">H", len(reply)) + reply) if not noise.handshake_finished: raise ProtocolError("handshake did not complete") return noise # --------------------------------------------------------------------------- # Operations (§6) # --------------------------------------------------------------------------- class Session: """Per-connection state: the authenticated username, if any.""" def __init__(self, server: "MailServer", peer_ip: str, handshake_hash: bytes) -> None: self.server = server self.peer_ip = peer_ip self.handshake_hash = handshake_hash self.username: str | None = None @property def store(self) -> Store: return self.server.store def dispatch(self, op: int, body: bytes) -> tuple[int, bytes]: handlers = { OP_AUTH: self.op_auth, OP_RESOLVE: self.op_resolve, OP_SEND: self.op_send, OP_FETCH: self.op_fetch, OP_DELETE: self.op_delete, OP_REGISTER: self.op_register, } handler = handlers.get(op) if handler is None: return MALFORMED, b"" if op in (OP_FETCH, OP_DELETE) and self.username is None: return AUTH_REQUIRED, b"" return handler(Reader(body)) def op_auth(self, r: Reader) -> tuple[int, bytes]: username = r.take(r.u8()).decode("utf-8", "strict") identity = r.take(KEY_LEN) signature = r.take(64) r.done() bound = self.store.identity_of(username) # A wrong username and a wrong signature are both AUTH_FAILED: telling # them apart would turn this into an account-existence oracle. if bound is None or bound != identity: return AUTH_FAILED, b"" if not verify(identity, signature, LABEL_AUTH + self.handshake_hash): return AUTH_FAILED, b"" self.username = username return OK, b"" def op_resolve(self, r: Reader) -> tuple[int, bytes]: username = r.take(r.u8()).decode("utf-8", "strict") r.done() identity = self.store.identity_of(username) if identity is None: return UNKNOWN_USER, b"" chain = self.store.chain(username) return OK, identity + bytes([len(chain)]) + b"".join(chain) def op_send(self, r: Reader) -> tuple[int, bytes]: envelope = r.rest() if not self.server.send_limiter.allow(self.peer_ip): return RATE_LIMITED, b"" if len(envelope) > self.server.max_envelope: return TOO_LARGE, b"" if len(envelope) < ENVELOPE_MIN or envelope[:4] != ENVELOPE_MAGIC: return MALFORMED, b"" if envelope[4] != ENVELOPE_VERSION: return BAD_VERSION, b"" recipient = envelope[5:37] username = self.store.username_for_key(recipient) if username is None: return UNKNOWN_USER, b"" keys = self.store.keys_of(username) if self.store.mailbox_bytes(keys) + len(envelope) > self.server.quota: return QUOTA_EXCEEDED, b"" # The ciphertext is never inspected; the server cannot read it. mid = message_id(envelope) self.store.store_message(mid, recipient, envelope) return OK, mid def op_fetch(self, r: Reader) -> tuple[int, bytes]: r.done() assert self.username is not None keys = self.store.keys_of(self.username) records = self.store.pending(keys, FETCH_BUDGET) out = [struct.pack(">H", len(records))] for mid, received_at, envelope in records: out.append(mid + struct.pack(">qI", received_at, len(envelope)) + envelope) return OK, b"".join(out) def op_delete(self, r: Reader) -> tuple[int, bytes]: count = r.u16() ids = [r.take(ID_LEN) for _ in range(count)] r.done() assert self.username is not None if not ids: return OK, struct.pack(">H", 0) # Scoped to the caller's own keys, so ids cannot be used to probe or # delete another mailbox. keys = self.store.keys_of(self.username) removed = self.store.delete(keys, ids) return OK, struct.pack(">H", removed) def op_register(self, r: Reader) -> tuple[int, bytes]: username = r.take(r.u8()).decode("utf-8", "strict") identity = r.take(KEY_LEN) token = r.take(r.u8()) cert = r.take(r.u8()) r.done() if not valid_username(username): return MALFORMED, b"" try: Ed25519PublicKey.from_public_bytes(identity) except ValueError: return MALFORMED, b"" expected = self.server.invite_token if expected is not None and token != expected: return NOT_PERMITTED, b"" if not cert: if not self.store.register(username, identity): return NOT_PERMITTED, b"" return OK, b"" if len(cert) != CERT_LEN: return MALFORMED, b"" old_pub, new_pub, when, signature = cert[:32], cert[32:64], cert[64:72], cert[72:] if new_pub != identity: return MALFORMED, b"" bound = self.store.identity_of(username) if bound is None: return UNKNOWN_USER, b"" if bound != old_pub: # Only the currently bound key may hand the username on (§7). return NOT_PERMITTED, b"" if not verify(old_pub, signature, LABEL_ROTATE + old_pub + new_pub + when): return AUTH_FAILED, b"" chain = self.store.chain(username) if len(chain) >= MAX_CHAIN: return NOT_PERMITTED, b"" if not self.store.rotate(username, new_pub, cert, len(chain)): return NOT_PERMITTED, b"" return OK, b"" # --------------------------------------------------------------------------- # Server # --------------------------------------------------------------------------- class Handler(socketserver.BaseRequestHandler): server: "MailServer" def handle(self) -> None: sock: socket.socket = self.request peer_ip = self.client_address[0] if not self.server.conn_limiter.allow(peer_ip): log.warning("rate limited %s", peer_ip) return sock.settimeout(IDLE_TIMEOUT) try: noise = handshake(sock, self.server.static_key) except Exception as exc: log.info("handshake failed from %s: %s", peer_ip, exc) return session = Session(self.server, peer_ip, noise.get_handshake_hash()) channel = Channel(sock, noise) while True: try: op, body = channel.read_frame() except (EOFError, OSError): return except ProtocolError as exc: log.info("bad frame from %s: %s", peer_ip, exc) self._try_send(channel, OP_AUTH, MALFORMED) return try: status, payload = session.dispatch(op, body) except (ProtocolError, UnicodeDecodeError) as exc: log.info("bad body from %s: %s", peer_ip, exc) self._try_send(channel, op, MALFORMED) return except Exception: log.exception("handler error from %s", peer_ip) self._try_send(channel, op, INTERNAL_ERROR) return try: channel.write_frame(op, bytes([status]) + payload) except OSError: return def _try_send(self, channel: Channel, op: int, status: int) -> None: try: channel.write_frame(op, bytes([status])) except OSError: pass class MailServer(socketserver.ThreadingTCPServer): allow_reuse_address = True daemon_threads = True def __init__(self, address, args: argparse.Namespace, static_key: bytes) -> None: super().__init__(address, Handler) self.static_key = static_key self.store = Store(args.db) self.max_envelope = args.max_envelope self.quota = args.quota self.retention = args.retention_days * 86400 self.invite_token = args.invite_token.encode() if args.invite_token else None self.conn_limiter = RateLimiter(args.rate_connections) self.send_limiter = RateLimiter(args.rate_sends) threading.Thread(target=self._purge_loop, daemon=True).start() def _purge_loop(self) -> None: store = Store(self.store.path) while True: time.sleep(PURGE_INTERVAL) try: gone = store.purge(int(time.time()) - self.retention) if gone: log.info("expired %d message(s)", gone) except Exception: log.exception("purge failed") # --------------------------------------------------------------------------- # CLI # --------------------------------------------------------------------------- def cmd_keygen(args: argparse.Namespace) -> int: if os.path.exists(args.key) and not args.force: print(f"{args.key} exists; refusing to overwrite (use --force)", file=sys.stderr) return 1 private = X25519PrivateKey.generate() raw = private.private_bytes( serialization.Encoding.Raw, serialization.PrivateFormat.Raw, serialization.NoEncryption(), ) # Written 0600 before any bytes land, so the key is never briefly readable. fd = os.open(args.key, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600) with os.fdopen(fd, "wb") as fh: fh.write(raw) public = private.public_key().public_bytes( serialization.Encoding.Raw, serialization.PublicFormat.Raw ) print(f"private key: {args.key}") print(f"public key: {b32(public)}") print() print("Publish the public key through a trusted channel; clients pin it (SPEC.md §4).") return 0 def cmd_serve(args: argparse.Namespace) -> int: try: with open(args.key, "rb") as fh: static_key = fh.read() except OSError as exc: print(f"cannot read server key: {exc}", file=sys.stderr) return 1 if len(static_key) != KEY_LEN: print(f"server key must be {KEY_LEN} raw bytes, got {len(static_key)}", file=sys.stderr) return 1 server = MailServer((args.host, args.port), args, static_key) public = ( X25519PrivateKey.from_private_bytes(static_key) .public_key() .public_bytes(serialization.Encoding.Raw, serialization.PublicFormat.Raw) ) log.info("listening on %s:%d", args.host, args.port) log.info("server public key: %s", b32(public)) try: server.serve_forever() except KeyboardInterrupt: log.info("shutting down") finally: server.shutdown() server.server_close() return 0 def main(argv: list[str] | None = None) -> int: parser = argparse.ArgumentParser(description=__doc__.splitlines()[0]) parser.add_argument("-v", "--verbose", action="store_true") sub = parser.add_subparsers(dest="command", required=True) keygen = sub.add_parser("keygen", help="generate the server's static X25519 key") keygen.add_argument("--key", default="server.key") keygen.add_argument("--force", action="store_true") keygen.set_defaults(func=cmd_keygen) serve = sub.add_parser("serve", help="run the mailbox server") serve.add_argument("--key", default="server.key") serve.add_argument("--db", default="mail.db") serve.add_argument("--host", default="127.0.0.1") serve.add_argument("--port", type=int, default=DEFAULT_PORT) serve.add_argument("--max-envelope", type=int, default=1 << 20, metavar="BYTES") serve.add_argument("--quota", type=int, default=64 << 20, metavar="BYTES") serve.add_argument("--retention-days", type=int, default=30) serve.add_argument("--invite-token", default=None, help="require this token to register; open registration if unset") serve.add_argument("--rate-connections", type=int, default=120, metavar="PER_MIN") serve.add_argument("--rate-sends", type=int, default=60, metavar="PER_MIN") serve.set_defaults(func=cmd_serve) args = parser.parse_args(argv) logging.basicConfig( level=logging.DEBUG if args.verbose else logging.INFO, format="%(asctime)s %(levelname)s %(message)s", ) return args.func(args) if __name__ == "__main__": sys.exit(main())