feat: revise to protocol 1.1 with accept tokens, fetch cursors and 32-byte ids

This commit is contained in:
randogoth 2026-09-27 08:27:08 +03:00
parent 56a6ed1186
commit 71f5f628be
4 changed files with 712 additions and 220 deletions

View file

@ -5,7 +5,7 @@
# ///
"""Smol Mail reference server.
Implements SPEC.md version 1: a store-and-forward mailbox reachable over TCP
Implements SPEC.md version 1.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.
@ -19,6 +19,7 @@ from __future__ import annotations
import argparse
import base64
import hashlib
import hmac
import logging
import os
import socket
@ -45,7 +46,9 @@ 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_MAC = b"smolmail/1 mac"
LABEL_ROTATE = b"smolmail/1 rotate"
LABEL_REGISTER = b"smolmail/1 register"
OP_AUTH = 0x00
OP_RESOLVE = 0x01
@ -70,16 +73,24 @@ 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
ID_LEN = 32
KEY_LEN = 32
CERT_LEN = 136 # old_pub 32 + new_pub 32 + time 8 + signature 64
SIG_LEN = 64
TOKEN_LEN = 32 # §5.8 accept token, and the MAC derived from it
CERT_LEN = 200 # old_pub 32 + new_pub 32 + time 8 + two signatures
MAX_CHAIN = 16
SEPARATORS = "._-"
TIER_MAIN = 0
TIER_REQUESTS = 1
FLAG_REQUESTS = 0x01 # §6.1 record flags
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
RECORD_OVERHEAD = ID_LEN + 8 + 1 + 4 # id + received_at + flags + env_len
IDLE_TIMEOUT = 120.0
PURGE_INTERVAL = 60.0
@ -137,17 +148,19 @@ class Reader:
def valid_username(name: str) -> bool:
"""§3: 1-63 bytes of [a-z0-9._-], not starting or ending with a separator."""
"""§3: 1-63 bytes of [a-z0-9._-], alphanumeric ends, no adjacent separators."""
if not 1 <= len(name) <= 63:
return False
if not all(c.isdigit() or ("a" <= c <= "z") or c in "._-" for c in name):
if not all("0" <= c <= "9" or "a" <= c <= "z" or c in SEPARATORS for c in name):
return False
return name[0] not in "._-" and name[-1] not in "._-"
if name[0] in SEPARATORS or name[-1] in SEPARATORS:
return False
return not any(a in SEPARATORS and b in SEPARATORS for a, b in zip(name, name[1:]))
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]
return hashlib.sha256(LABEL_ID + envelope).digest()
def verify(pubkey: bytes, signature: bytes, message: bytes) -> bool:
@ -179,14 +192,30 @@ CREATE TABLE IF NOT EXISTS rotations (
cert BLOB NOT NULL,
PRIMARY KEY (username, seq)
);
-- Accept tokens (§5.8): opaque secrets the mailbox owner uploads with AUTH.
-- The server can match them against a message's MAC and nothing else; it never
-- learns which correspondent a token belongs to.
CREATE TABLE IF NOT EXISTS accepted (
username TEXT NOT NULL,
token BLOB NOT NULL,
PRIMARY KEY (username, token)
);
CREATE TABLE IF NOT EXISTS messages (
id BLOB PRIMARY KEY,
recipient BLOB NOT NULL,
received_at INTEGER NOT NULL,
tier INTEGER NOT NULL,
envelope BLOB NOT NULL
);
-- (received_at, id) is FETCH's sort and cursor order (§6.1).
CREATE INDEX IF NOT EXISTS messages_by_recipient
ON messages (recipient, received_at);
ON messages (recipient, received_at, id);
-- Identifiers of every envelope accepted within the retention window, so an
-- envelope captured and resent after DELETE does not reappear (§10).
CREATE TABLE IF NOT EXISTS seen (
id BLOB PRIMARY KEY,
at INTEGER NOT NULL
);
"""
@ -277,39 +306,78 @@ class Store:
except sqlite3.IntegrityError:
return False
def mailbox_bytes(self, keys: list[bytes]) -> int:
def accepted_tokens(self, username: str) -> list[bytes]:
return [
row[0]
for row in self.db.execute(
"SELECT token FROM accepted WHERE username = ?", (username,)
)
]
def set_accepted(self, username: str, tokens: list[bytes]) -> int:
"""§4: an AUTH with sync = 1 replaces the whole set, which is how a
token is both added and removed."""
with self.db:
self.db.execute("DELETE FROM accepted WHERE username = ?", (username,))
self.db.executemany(
"INSERT OR IGNORE INTO accepted (username, token) VALUES (?, ?)",
[(username, token) for token in tokens],
)
return self.count_accepted(username)
def count_accepted(self, username: str) -> int:
return self.db.execute(
"SELECT COUNT(*) FROM accepted WHERE username = ?", (username,)
).fetchone()[0]
def tombstoned(self, mid: bytes) -> bool:
return (
self.db.execute("SELECT 1 FROM seen WHERE id = ?", (mid,)).fetchone()
is not None
)
def mailbox_bytes(self, keys: list[bytes], tier: int) -> int:
marks = ",".join("?" * len(keys))
row = self.db.execute(
f"SELECT COALESCE(SUM(LENGTH(envelope)), 0) FROM messages "
f"WHERE recipient IN ({marks})",
keys,
f"WHERE tier = ? AND recipient IN ({marks})",
[tier] + keys,
).fetchone()
return row[0]
def store_message(self, mid: bytes, recipient: bytes, envelope: bytes) -> None:
def store_message(self, mid: bytes, recipient: bytes, tier: int, envelope: bytes) -> None:
with self.db:
now = int(time.time())
self.db.execute(
"INSERT OR IGNORE INTO messages (id, recipient, received_at, envelope) "
"VALUES (?, ?, ?, ?)",
(mid, recipient, int(time.time()), envelope),
"INSERT OR IGNORE INTO messages "
"(id, recipient, received_at, tier, envelope) VALUES (?, ?, ?, ?, ?)",
(mid, recipient, now, tier, envelope),
)
self.db.execute(
"INSERT OR IGNORE INTO seen (id, at) VALUES (?, ?)", (mid, now)
)
def pending(self, keys: list[bytes], budget: int) -> list[tuple[bytes, int, bytes]]:
def pending(
self, keys: list[bytes], after_time: int, after_id: bytes, budget: int
) -> list[tuple[bytes, int, 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,
f"SELECT id, received_at, tier, envelope FROM messages "
f"WHERE recipient IN ({marks}) "
f"AND (received_at > ? OR (received_at = ? AND id > ?)) "
f"ORDER BY received_at, id",
keys + [after_time, after_time, after_id],
)
out: list[tuple[bytes, int, bytes]] = []
out: list[tuple[bytes, int, int, bytes]] = []
used = 0
for mid, received_at, envelope in rows:
# Always return at least one message, even if it alone exceeds the
for mid, received_at, tier, envelope in rows:
size = RECORD_OVERHEAD + len(envelope)
# Always return at least one record, even if it alone exceeds the
# budget, so an oversized envelope cannot wedge a mailbox shut.
if out and used + len(envelope) > budget:
if out and used + size > budget:
break
out.append((mid, received_at, envelope))
used += len(envelope)
out.append((mid, received_at, tier, envelope))
used += size
return out
def delete(self, keys: list[bytes], ids: list[bytes]) -> int:
@ -323,10 +391,16 @@ class Store:
)
return cur.rowcount
def purge(self, older_than: int) -> int:
def purge(self, main_cutoff: int, requests_cutoff: int, seen_cutoff: int) -> int:
with self.db:
cur = self.db.execute("DELETE FROM messages WHERE received_at < ?", (older_than,))
return cur.rowcount
cur = self.db.execute(
"DELETE FROM messages WHERE (tier = ? AND received_at < ?) "
"OR (tier = ? AND received_at < ?)",
(TIER_MAIN, main_cutoff, TIER_REQUESTS, requests_cutoff),
)
gone = cur.rowcount
self.db.execute("DELETE FROM seen WHERE at < ?", (seen_cutoff,))
return gone
# ---------------------------------------------------------------------------
@ -335,10 +409,7 @@ class Store:
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.
"""
"""Fixed-window counter, keyed by IP address or by accept token (§10)."""
def __init__(self, limit: int, window: float = 60.0) -> None:
self.limit = limit
@ -346,25 +417,25 @@ class RateLimiter:
self.lock = threading.Lock()
self.hits: dict[str, tuple[float, int]] = {}
def allow(self, ip: str) -> bool:
def allow(self, key: str) -> bool:
if self.limit <= 0:
return True
now = time.monotonic()
with self.lock:
start, count = self.hits.get(ip, (now, 0))
start, count = self.hits.get(key, (now, 0))
if now - start >= self.window:
start, count = now, 0
if count >= self.limit:
return False
self.hits[ip] = (start, count + 1)
self.hits[key] = (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]
stale = [key for key, (start, _) in self.hits.items() if now - start >= self.window]
for key in stale:
del self.hits[key]
# ---------------------------------------------------------------------------
@ -475,7 +546,9 @@ class Session:
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)
signature = r.take(SIG_LEN)
sync = r.u8()
tokens = [r.take(TOKEN_LEN) for _ in range(r.u16())]
r.done()
bound = self.store.identity_of(username)
# A wrong username and a wrong signature are both AUTH_FAILED: telling
@ -484,8 +557,16 @@ class Session:
return AUTH_FAILED, b""
if not verify(identity, signature, LABEL_AUTH + self.handshake_hash):
return AUTH_FAILED, b""
if sync not in (0, 1) or (sync == 0 and tokens):
return MALFORMED, b""
# Refused before authenticating, so the client has to trim and retry
# rather than silently running with a truncated set (§4).
if len(tokens) > self.server.max_accepted:
return TOO_LARGE, b""
self.username = username
return OK, b""
held = self.store.set_accepted(username, tokens) if sync else \
self.store.count_accepted(username)
return OK, struct.pack(">H", held)
def op_resolve(self, r: Reader) -> tuple[int, bytes]:
username = r.take(r.u8()).decode("utf-8", "strict")
@ -496,7 +577,22 @@ class Session:
chain = self.store.chain(username)
return OK, identity + bytes([len(chain)]) + b"".join(chain)
def match_token(self, username: str, mid: bytes, mac: bytes) -> bytes | None:
"""§5.8. Trial-match the MAC against the mailbox's tokens.
Bounded by --max-accepted, so an unmatchable MAC costs a known amount
of HMAC rather than an unbounded scan.
"""
if len(mac) != TOKEN_LEN:
return None
for token in self.store.accepted_tokens(username):
want = hmac.new(token, LABEL_MAC + mid, hashlib.sha256).digest()
if hmac.compare_digest(want, mac):
return token
return None
def op_send(self, r: Reader) -> tuple[int, bytes]:
mac = r.take(r.u8())
envelope = r.rest()
if not self.server.send_limiter.allow(self.peer_ip):
return RATE_LIMITED, b""
@ -510,22 +606,40 @@ class Session:
username = self.store.username_for_key(recipient)
if username is None:
return UNKNOWN_USER, b""
mid = message_id(envelope)
# A resend is answered with the same id and stores nothing, whether the
# original is still here or was deleted (§10).
if self.store.tombstoned(mid):
return OK, mid
# An unmatched MAC is not an error: SEND must not reveal whether a
# token is still accepted (§5.8).
token = self.match_token(username, mid, mac)
if token is None:
tier, quota = TIER_REQUESTS, self.server.requests_quota
else:
tier, quota = TIER_MAIN, self.server.quota
if not self.server.token_limiter.allow(token.hex()):
return RATE_LIMITED, b""
keys = self.store.keys_of(username)
if self.store.mailbox_bytes(keys) + len(envelope) > self.server.quota:
if self.store.mailbox_bytes(keys, tier) + len(envelope) > 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)
self.store.store_message(mid, recipient, tier, envelope)
return OK, mid
def op_fetch(self, r: Reader) -> tuple[int, bytes]:
after_time = r.i64()
after_id = r.take(ID_LEN)
r.done()
assert self.username is not None
keys = self.store.keys_of(self.username)
records = self.store.pending(keys, FETCH_BUDGET)
records = self.store.pending(keys, after_time, after_id, 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)
for mid, received_at, tier, envelope in records:
flags = FLAG_REQUESTS if tier == TIER_REQUESTS else 0
out.append(
mid + struct.pack(">qBI", received_at, flags, len(envelope)) + envelope
)
return OK, b"".join(out)
def op_delete(self, r: Reader) -> tuple[int, bytes]:
@ -544,6 +658,7 @@ class Session:
def op_register(self, r: Reader) -> tuple[int, bytes]:
username = r.take(r.u8()).decode("utf-8", "strict")
identity = r.take(KEY_LEN)
signature = r.take(SIG_LEN)
token = r.take(r.u8())
cert = r.take(r.u8())
r.done()
@ -554,6 +669,14 @@ class Session:
Ed25519PublicKey.from_public_bytes(identity)
except ValueError:
return MALFORMED, b""
# §6.1 proof of possession, bound to this server's static key so the
# attestation cannot be replayed elsewhere.
if not verify(
identity,
signature,
LABEL_REGISTER + self.server.static_pub + username.encode() + identity,
):
return AUTH_FAILED, b""
expected = self.server.invite_token
if expected is not None and token != expected:
return NOT_PERMITTED, b""
@ -565,7 +688,8 @@ class Session:
if len(cert) != CERT_LEN:
return MALFORMED, b""
old_pub, new_pub, when, signature = cert[:32], cert[32:64], cert[64:72], cert[72:]
old_pub, new_pub, when = cert[:32], cert[32:64], cert[64:72]
sig_old, sig_new = cert[72:136], cert[136:200]
if new_pub != identity:
return MALFORMED, b""
bound = self.store.identity_of(username)
@ -574,7 +698,10 @@ class Session:
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):
# Both keys sign: the old one alone could otherwise rotate to a key
# nobody controls (§7).
signed = LABEL_ROTATE + username.encode() + old_pub + new_pub + when
if not verify(old_pub, sig_old, signed) or not verify(new_pub, sig_new, signed):
return AUTH_FAILED, b""
chain = self.store.chain(username)
if len(chain) >= MAX_CHAIN:
@ -638,6 +765,14 @@ class Handler(socketserver.BaseRequestHandler):
pass
def static_public(static_key: bytes) -> bytes:
return (
X25519PrivateKey.from_private_bytes(static_key)
.public_key()
.public_bytes(serialization.Encoding.Raw, serialization.PublicFormat.Raw)
)
class MailServer(socketserver.ThreadingTCPServer):
allow_reuse_address = True
daemon_threads = True
@ -645,13 +780,18 @@ class MailServer(socketserver.ThreadingTCPServer):
def __init__(self, address, args: argparse.Namespace, static_key: bytes) -> None:
super().__init__(address, Handler)
self.static_key = static_key
self.static_pub = static_public(static_key)
self.store = Store(args.db)
self.max_envelope = args.max_envelope
self.quota = args.quota
self.requests_quota = args.requests_quota
self.retention = args.retention_days * 86400
self.requests_retention = args.requests_retention_days * 86400
self.max_accepted = args.max_accepted
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)
self.token_limiter = RateLimiter(args.rate_token)
threading.Thread(target=self._purge_loop, daemon=True).start()
def _purge_loop(self) -> None:
@ -659,7 +799,14 @@ class MailServer(socketserver.ThreadingTCPServer):
while True:
time.sleep(PURGE_INTERVAL)
try:
gone = store.purge(int(time.time()) - self.retention)
now = int(time.time())
# Tombstones outlive both tiers, so a replay cannot slip in
# between a message expiring and its id being forgotten.
gone = store.purge(
now - self.retention,
now - self.requests_retention,
now - max(self.retention, self.requests_retention),
)
if gone:
log.info("expired %d message(s)", gone)
except Exception:
@ -707,13 +854,8 @@ def cmd_serve(args: argparse.Namespace) -> int:
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))
log.info("server public key: %s", b32(server.static_pub))
try:
server.serve_forever()
except KeyboardInterrupt:
@ -739,13 +881,20 @@ def main(argv: list[str] | None = None) -> int:
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("--max-envelope", type=int, default=768 << 10, metavar="BYTES")
serve.add_argument("--quota", type=int, default=64 << 20, metavar="BYTES")
serve.add_argument("--requests-quota", type=int, default=2 << 20, metavar="BYTES",
help="quota for mail arriving without an accept token")
serve.add_argument("--retention-days", type=int, default=30)
serve.add_argument("--requests-retention-days", type=int, default=7)
serve.add_argument("--max-accepted", type=int, default=1024, metavar="TOKENS",
help="accept tokens a mailbox may hold (SPEC.md §5.8)")
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.add_argument("--rate-token", type=int, default=120, metavar="PER_MIN",
help="sends per minute per accept token")
serve.set_defaults(func=cmd_serve)
args = parser.parse_args(argv)