911 lines
33 KiB
Python
Executable file
911 lines
33 KiB
Python
Executable file
#!/usr/bin/env -S uv run --quiet --script
|
|
# /// script
|
|
# requires-python = ">=3.11"
|
|
# dependencies = ["noiseprotocol>=0.3.1", "cryptography>=42"]
|
|
# ///
|
|
# Copyright 2026 randogoth
|
|
# SPDX-License-Identifier: Apache-2.0
|
|
"""Smol Mail reference server.
|
|
|
|
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.
|
|
|
|
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 hmac
|
|
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_MAC = b"smolmail/1 mac"
|
|
LABEL_ROTATE = b"smolmail/1 rotate"
|
|
LABEL_REGISTER = b"smolmail/1 register"
|
|
|
|
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 = 32
|
|
KEY_LEN = 32
|
|
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
|
|
|
|
|
|
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._-], alphanumeric ends, no adjacent separators."""
|
|
if not 1 <= len(name) <= 63:
|
|
return False
|
|
if not all("0" <= c <= "9" or "a" <= c <= "z" or c in SEPARATORS for c in name):
|
|
return False
|
|
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()
|
|
|
|
|
|
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)
|
|
);
|
|
-- 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, 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
|
|
);
|
|
"""
|
|
|
|
|
|
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 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 tier = ? AND recipient IN ({marks})",
|
|
[tier] + keys,
|
|
).fetchone()
|
|
return row[0]
|
|
|
|
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, 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], 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, 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, int, bytes]] = []
|
|
used = 0
|
|
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 + size > budget:
|
|
break
|
|
out.append((mid, received_at, tier, envelope))
|
|
used += size
|
|
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, main_cutoff: int, requests_cutoff: int, seen_cutoff: int) -> int:
|
|
with self.db:
|
|
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
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Rate limiting
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class RateLimiter:
|
|
"""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
|
|
self.window = window
|
|
self.lock = threading.Lock()
|
|
self.hits: dict[str, tuple[float, int]] = {}
|
|
|
|
def allow(self, key: str) -> bool:
|
|
if self.limit <= 0:
|
|
return True
|
|
now = time.monotonic()
|
|
with self.lock:
|
|
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[key] = (start, count + 1)
|
|
if len(self.hits) > 4096:
|
|
self._evict(now)
|
|
return True
|
|
|
|
def _evict(self, now: float) -> None:
|
|
stale = [key for key, (start, _) in self.hits.items() if now - start >= self.window]
|
|
for key in stale:
|
|
del self.hits[key]
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# 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(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
|
|
# 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""
|
|
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
|
|
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")
|
|
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 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""
|
|
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""
|
|
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, tier) + len(envelope) > quota:
|
|
return QUOTA_EXCEEDED, b""
|
|
# The ciphertext is never inspected; the server cannot read it.
|
|
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, after_time, after_id, FETCH_BUDGET)
|
|
out = [struct.pack(">H", len(records))]
|
|
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]:
|
|
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)
|
|
signature = r.take(SIG_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""
|
|
# §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""
|
|
|
|
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 = 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)
|
|
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""
|
|
# 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:
|
|
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
|
|
|
|
|
|
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
|
|
|
|
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:
|
|
store = Store(self.store.path)
|
|
while True:
|
|
time.sleep(PURGE_INTERVAL)
|
|
try:
|
|
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:
|
|
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)
|
|
log.info("listening on %s:%d", args.host, args.port)
|
|
log.info("server public key: %s", b32(server.static_pub))
|
|
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=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)
|
|
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())
|