feat: add single-file reference server
This commit is contained in:
parent
9a31607648
commit
69ba092179
2 changed files with 773 additions and 0 deletions
760
smolmaild.py
Executable file
760
smolmaild.py
Executable file
|
|
@ -0,0 +1,760 @@
|
|||
#!/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())
|
||||
Loading…
Add table
Add a link
Reference in a new issue