feat: implement Smol Mail server in Rust with nix flake deployment

This commit is contained in:
randogoth 2026-09-26 11:52:22 +03:00
commit 71ad7e04c0
14 changed files with 2592 additions and 0 deletions

122
src/channel.rs Normal file
View file

@ -0,0 +1,122 @@
//! Noise_NX handshake and the framed transport on top of it.
//!
//! Two independent layers: Noise messages under a u16 length prefix, and
//! application frames (u32 length || u8 op || body) split across as many
//! Noise messages as they need and reassembled from them.
use std::io::{self, Read, Write};
use std::net::TcpStream;
use snow::{Builder, TransportState};
use crate::proto::{ProtocolError, MAX_FRAME, NOISE_PARAMS, NOISE_PAYLOAD, PROLOGUE};
/// Runs the Noise_NX responder handshake. The initiator stays anonymous;
/// only we hold a static key. Returns the transport state and the
/// handshake hash (needed later to verify AUTH frames).
pub fn handshake(
stream: &mut TcpStream,
static_key: &[u8],
) -> anyhow::Result<(TransportState, Vec<u8>)> {
let params: snow::params::NoiseParams = NOISE_PARAMS.parse()?;
let mut noise = Builder::new(params)
.local_private_key(static_key)
.prologue(PROLOGUE)
.build_responder()?;
let mut buf = [0u8; 65535];
let mut msg = [0u8; 65535];
let len = read_u16_len(stream)?;
read_exact_into(stream, &mut buf[..len])?;
noise.read_message(&buf[..len], &mut msg)?;
let len = noise.write_message(&[], &mut buf)?;
write_u16_len(stream, &buf[..len])?;
anyhow::ensure!(noise.is_handshake_finished(), "handshake did not complete");
let hash = noise.get_handshake_hash().to_vec();
let transport = noise.into_transport_mode()?;
Ok((transport, hash))
}
fn read_exact_into(stream: &mut TcpStream, buf: &mut [u8]) -> io::Result<()> {
stream.read_exact(buf)
}
fn read_u16_len(stream: &mut TcpStream) -> io::Result<usize> {
let mut len_buf = [0u8; 2];
stream.read_exact(&mut len_buf)?;
Ok(u16::from_be_bytes(len_buf) as usize)
}
fn write_u16_len(stream: &mut TcpStream, packet: &[u8]) -> io::Result<()> {
stream.write_all(&(packet.len() as u16).to_be_bytes())?;
stream.write_all(packet)?;
Ok(())
}
pub struct Channel {
stream: TcpStream,
transport: TransportState,
buf: Vec<u8>,
}
impl Channel {
pub fn new(stream: TcpStream, transport: TransportState) -> Self {
Channel {
stream,
transport,
buf: Vec::new(),
}
}
fn read_noise(&mut self) -> anyhow::Result<Vec<u8>> {
let len = read_u16_len(&mut self.stream)?;
let mut ciphertext = vec![0u8; len];
read_exact_into(&mut self.stream, &mut ciphertext)?;
let mut plaintext = vec![0u8; len];
let n = self.transport.read_message(&ciphertext, &mut plaintext)?;
plaintext.truncate(n);
Ok(plaintext)
}
fn write_noise(&mut self, payload: &[u8]) -> anyhow::Result<()> {
let mut packet = vec![0u8; payload.len() + 16];
let n = self.transport.write_message(payload, &mut packet)?;
packet.truncate(n);
write_u16_len(&mut self.stream, &packet)?;
Ok(())
}
/// Reads one application frame, blocking until a full frame is available.
pub fn read_frame(&mut self) -> anyhow::Result<(u8, Vec<u8>)> {
while self.buf.len() < 5 {
let chunk = self.read_noise()?;
self.buf.extend_from_slice(&chunk);
}
let length = u32::from_be_bytes(self.buf[..4].try_into().unwrap()) as usize;
if length < 1 || length > MAX_FRAME {
return Err(ProtocolError::new(format!("frame length {length} out of range")).into());
}
while self.buf.len() < 4 + length {
let chunk = self.read_noise()?;
self.buf.extend_from_slice(&chunk);
}
let frame: Vec<u8> = self.buf[4..4 + length].to_vec();
self.buf.drain(..4 + length);
Ok((frame[0], frame[1..].to_vec()))
}
pub fn write_frame(&mut self, op: u8, body: &[u8]) -> anyhow::Result<()> {
let mut frame = Vec::with_capacity(5 + body.len());
frame.extend_from_slice(&((1 + body.len()) as u32).to_be_bytes());
frame.push(op);
frame.extend_from_slice(body);
for chunk in frame.chunks(NOISE_PAYLOAD) {
self.write_noise(chunk)?;
}
Ok(())
}
}

88
src/crypto.rs Normal file
View file

@ -0,0 +1,88 @@
//! Encoding and cryptographic verification helpers.
use data_encoding::{Encoding, Specification};
use ed25519_dalek::{Signature, VerifyingKey};
use rand_core::OsRng;
use sha2::{Digest, Sha256};
use std::sync::LazyLock;
use x25519_dalek::{PublicKey, StaticSecret};
use crate::proto::{LABEL_ID, ID_LEN};
static BASE32_LOWER_UNPADDED: LazyLock<Encoding> = LazyLock::new(|| {
let mut spec = Specification::new();
spec.symbols.push_str("abcdefghijklmnopqrstuvwxyz234567");
spec.encoding().unwrap()
});
/// RFC 4648 base32, lowercase and unpadded.
pub fn b32(raw: &[u8]) -> String {
BASE32_LOWER_UNPADDED.encode(raw)
}
/// Derived from the envelope, so a sender cannot choose it.
pub fn message_id(envelope: &[u8]) -> [u8; ID_LEN] {
let mut hasher = Sha256::new();
hasher.update(LABEL_ID);
hasher.update(envelope);
let digest = hasher.finalize();
let mut out = [0u8; ID_LEN];
out.copy_from_slice(&digest[..ID_LEN]);
out
}
/// Generates the server's static X25519 keypair (transport identity, distinct
/// from any user's Ed25519 identity). Returns (private, public) raw bytes.
pub fn generate_static_key() -> ([u8; 32], [u8; 32]) {
let secret = StaticSecret::random_from_rng(OsRng);
let public = PublicKey::from(&secret);
(secret.to_bytes(), public.to_bytes())
}
/// Derives the X25519 public key for a raw static private key.
pub fn derive_public(private: &[u8; 32]) -> [u8; 32] {
let secret = StaticSecret::from(*private);
PublicKey::from(&secret).to_bytes()
}
/// Verifies an Ed25519 signature; malformed keys/signatures are simply not valid.
pub fn verify(pubkey: &[u8], signature: &[u8], message: &[u8]) -> bool {
let Ok(pubkey): Result<[u8; 32], _> = pubkey.try_into() else {
return false;
};
let Ok(signature): Result<[u8; 64], _> = signature.try_into() else {
return false;
};
let Ok(verifying_key) = VerifyingKey::from_bytes(&pubkey) else {
return false;
};
let signature = Signature::from_bytes(&signature);
verifying_key.verify_strict(message, &signature).is_ok()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn base32_matches_rfc4648_lowercase_unpadded() {
// "hello" -> base32 "NBSWY3DP" per RFC 4648, lowercased and unpadded.
assert_eq!(b32(b"hello"), "nbswy3dp");
}
#[test]
fn message_id_is_16_bytes_and_deterministic() {
let a = message_id(b"envelope-bytes");
let b = message_id(b"envelope-bytes");
let c = message_id(b"other-bytes");
assert_eq!(a.len(), ID_LEN);
assert_eq!(a, b);
assert_ne!(a, c);
}
#[test]
fn verify_rejects_garbage() {
assert!(!verify(&[0u8; 32], &[0u8; 64], b"msg"));
assert!(!verify(&[0u8; 5], &[0u8; 64], b"msg"));
}
}

123
src/main.rs Normal file
View file

@ -0,0 +1,123 @@
//! Smol Mail server.
mod channel;
mod crypto;
mod proto;
mod ratelimit;
mod server;
mod session;
mod store;
use std::io::Write;
#[cfg(unix)]
use std::os::unix::fs::OpenOptionsExt;
use clap::{Parser, Subcommand};
use crate::proto::DEFAULT_PORT;
#[derive(Parser)]
#[command(about = "Smol Mail server")]
struct Cli {
#[arg(short, long, global = true)]
verbose: bool,
#[command(subcommand)]
command: Command,
}
#[derive(Subcommand)]
enum Command {
/// Generate the server's static X25519 key
Keygen {
#[arg(long, default_value = "server.key")]
key: String,
#[arg(long)]
force: bool,
},
/// Run the mailbox server
Serve {
#[arg(long, default_value = "server.key")]
key: String,
#[arg(long, default_value = "mail.db")]
db: String,
#[arg(long, default_value = "127.0.0.1")]
host: String,
#[arg(long, default_value_t = DEFAULT_PORT)]
port: u16,
#[arg(long = "max-envelope", default_value_t = 1 << 20)]
max_envelope: usize,
#[arg(long, default_value_t = 64 << 20)]
quota: i64,
#[arg(long = "retention-days", default_value_t = 30)]
retention_days: i64,
#[arg(long = "invite-token")]
invite_token: Option<String>,
#[arg(long = "rate-connections", default_value_t = 120)]
rate_connections: u32,
#[arg(long = "rate-sends", default_value_t = 60)]
rate_sends: u32,
},
}
fn main() -> anyhow::Result<()> {
let cli = Cli::parse();
env_logger::Builder::new()
.filter_level(if cli.verbose {
log::LevelFilter::Debug
} else {
log::LevelFilter::Info
})
.format_timestamp_secs()
.init();
match cli.command {
Command::Keygen { key, force } => cmd_keygen(&key, force),
Command::Serve {
key,
db,
host,
port,
max_envelope,
quota,
retention_days,
invite_token,
rate_connections,
rate_sends,
} => server::run(server::ServeArgs {
key_path: key,
db_path: db,
host,
port,
max_envelope,
quota,
retention_days,
invite_token,
rate_connections,
rate_sends,
}),
}
}
fn cmd_keygen(key_path: &str, force: bool) -> anyhow::Result<()> {
if std::path::Path::new(key_path).exists() && !force {
anyhow::bail!("{key_path} exists; refusing to overwrite (use --force)");
}
let (private, public) = crypto::generate_static_key();
// Written 0600 before any bytes land, so the key is never briefly readable.
let mut opts = std::fs::OpenOptions::new();
opts.write(true).create(true).truncate(true);
#[cfg(unix)]
opts.mode(0o600);
let mut file = opts.open(key_path)?;
file.write_all(&private)?;
println!("private key: {key_path}");
println!("public key: {}", crypto::b32(&public));
println!();
println!("Publish the public key through a trusted channel; clients pin it (SPEC.md sec 4).");
Ok(())
}

166
src/proto.rs Normal file
View file

@ -0,0 +1,166 @@
//! Wire constants and body parsing.
use std::fmt;
pub const NOISE_PARAMS: &str = "Noise_NX_25519_ChaChaPoly_SHA256";
pub const PROLOGUE: &[u8] = b"smolmail/1";
pub const LABEL_AUTH: &[u8] = b"smolmail/1 auth";
pub const LABEL_ID: &[u8] = b"smolmail/1 id";
pub const LABEL_ROTATE: &[u8] = b"smolmail/1 rotate";
pub const OP_AUTH: u8 = 0x00;
pub const OP_RESOLVE: u8 = 0x01;
pub const OP_SEND: u8 = 0x02;
pub const OP_FETCH: u8 = 0x03;
pub const OP_DELETE: u8 = 0x04;
pub const OP_REGISTER: u8 = 0x05;
pub const OK: u8 = 0;
pub const MALFORMED: u8 = 1;
pub const BAD_VERSION: u8 = 2;
pub const UNKNOWN_USER: u8 = 3;
pub const AUTH_REQUIRED: u8 = 4;
pub const AUTH_FAILED: u8 = 5;
pub const QUOTA_EXCEEDED: u8 = 6;
pub const TOO_LARGE: u8 = 7;
pub const RATE_LIMITED: u8 = 8;
pub const NOT_PERMITTED: u8 = 9;
pub const INTERNAL_ERROR: u8 = 10;
pub const ENVELOPE_MAGIC: &[u8; 4] = b"SMOL";
pub const ENVELOPE_VERSION: u8 = 1;
pub const ENVELOPE_HEADER: usize = 69; // magic 4 + version 1 + to 32 + epk 32
pub const ENVELOPE_MIN: usize = ENVELOPE_HEADER + 16; // + Poly1305 tag
pub const ID_LEN: usize = 16;
pub const KEY_LEN: usize = 32;
pub const CERT_LEN: usize = 136; // old_pub 32 + new_pub 32 + time 8 + signature 64
pub const MAX_CHAIN: usize = 16;
pub const DEFAULT_PORT: u16 = 1961;
pub const MAX_FRAME: usize = 1 << 20; // application frame ceiling
pub const NOISE_MAX: usize = 65535; // Noise message ceiling
pub const NOISE_PAYLOAD: usize = NOISE_MAX - 16; // minus the AEAD tag
pub const FETCH_BUDGET: usize = 512 * 1024; // must stay under MAX_FRAME
pub const IDLE_TIMEOUT_SECS: u64 = 120;
pub const PURGE_INTERVAL_SECS: u64 = 60;
/// A peer sent something unparseable. Always answered with MALFORMED.
#[derive(Debug)]
pub struct ProtocolError(pub String);
impl fmt::Display for ProtocolError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}", self.0)
}
}
impl std::error::Error for ProtocolError {}
impl ProtocolError {
pub fn new(msg: impl Into<String>) -> Self {
ProtocolError(msg.into())
}
}
/// Fail-closed reader over a frame body.
///
/// Every parse path errors rather than reading past the end, so a truncated
/// frame can never be mistaken for a short but valid one.
pub struct Reader<'a> {
buf: &'a [u8],
pos: usize,
}
impl<'a> Reader<'a> {
pub fn new(buf: &'a [u8]) -> Self {
Reader { buf, pos: 0 }
}
pub fn take(&mut self, n: usize) -> Result<&'a [u8], ProtocolError> {
if self.pos + n > self.buf.len() {
return Err(ProtocolError::new(format!(
"short read: want {}, have {}",
n,
self.buf.len() - self.pos
)));
}
let out = &self.buf[self.pos..self.pos + n];
self.pos += n;
Ok(out)
}
pub fn u8(&mut self) -> Result<u8, ProtocolError> {
Ok(self.take(1)?[0])
}
pub fn u16(&mut self) -> Result<u16, ProtocolError> {
let b = self.take(2)?;
Ok(u16::from_be_bytes([b[0], b[1]]))
}
pub fn rest(&mut self) -> &'a [u8] {
let out = &self.buf[self.pos..];
self.pos = self.buf.len();
out
}
pub fn done(&self) -> Result<(), ProtocolError> {
if self.pos != self.buf.len() {
return Err(ProtocolError::new(format!(
"{} trailing bytes",
self.buf.len() - self.pos
)));
}
Ok(())
}
}
/// 1-63 bytes of [a-z0-9._-], not starting or ending with a separator.
pub fn valid_username(name: &str) -> bool {
let bytes = name.as_bytes();
if bytes.is_empty() || bytes.len() > 63 {
return false;
}
if !bytes
.iter()
.all(|&c| c.is_ascii_digit() || c.is_ascii_lowercase() || matches!(c, b'.' | b'_' | b'-'))
{
return false;
}
let first = bytes[0];
let last = bytes[bytes.len() - 1];
!matches!(first, b'.' | b'_' | b'-') && !matches!(last, b'.' | b'_' | b'-')
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn reader_bounds() {
let mut r = Reader::new(&[1, 2, 3]);
assert_eq!(r.u8().unwrap(), 1);
assert!(r.take(3).is_err());
assert_eq!(r.take(2).unwrap(), &[2, 3]);
assert!(r.done().is_ok());
}
#[test]
fn reader_trailing_bytes_rejected() {
let mut r = Reader::new(&[1, 2, 3]);
let _ = r.u8().unwrap();
assert!(r.done().is_err());
}
#[test]
fn username_validation() {
assert!(valid_username("alice"));
assert!(valid_username("a.b_c-9"));
assert!(!valid_username(""));
assert!(!valid_username(&"a".repeat(64)));
assert!(!valid_username(".alice"));
assert!(!valid_username("alice."));
assert!(!valid_username("Alice"));
assert!(!valid_username("al ice"));
}
}

82
src/ratelimit.rs Normal file
View file

@ -0,0 +1,82 @@
//! Fixed-window per-IP counter, the whole of the server's abuse control.
//!
//! A server cannot see senders, so quotas, size caps and this are all it has.
use std::collections::HashMap;
use std::sync::Mutex;
use std::time::{Duration, Instant};
struct Window {
start: Instant,
count: u32,
}
pub struct RateLimiter {
limit: u32,
window: Duration,
hits: Mutex<HashMap<String, Window>>,
}
impl RateLimiter {
pub fn new(limit: u32) -> Self {
RateLimiter {
limit,
window: Duration::from_secs(60),
hits: Mutex::new(HashMap::new()),
}
}
pub fn allow(&self, ip: &str) -> bool {
if self.limit == 0 {
return true;
}
let now = Instant::now();
let mut hits = self.hits.lock().unwrap();
let entry = hits.entry(ip.to_string()).or_insert(Window {
start: now,
count: 0,
});
if now.duration_since(entry.start) >= self.window {
entry.start = now;
entry.count = 0;
}
if entry.count >= self.limit {
return false;
}
entry.count += 1;
if hits.len() > 4096 {
let window = self.window;
hits.retain(|_, w| now.duration_since(w.start) < window);
}
true
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn allows_up_to_limit_then_blocks() {
let rl = RateLimiter::new(2);
assert!(rl.allow("1.2.3.4"));
assert!(rl.allow("1.2.3.4"));
assert!(!rl.allow("1.2.3.4"));
}
#[test]
fn zero_limit_means_unlimited() {
let rl = RateLimiter::new(0);
for _ in 0..100 {
assert!(rl.allow("1.2.3.4"));
}
}
#[test]
fn separate_ips_have_separate_windows() {
let rl = RateLimiter::new(1);
assert!(rl.allow("1.1.1.1"));
assert!(rl.allow("2.2.2.2"));
assert!(!rl.allow("1.1.1.1"));
}
}

163
src/server.rs Normal file
View file

@ -0,0 +1,163 @@
//! TCP accept loop, per-connection handling, and the background purge loop.
use std::net::TcpStream;
use std::sync::Arc;
use std::time::Duration;
use crate::channel::{handshake, Channel};
use crate::crypto::{b32, derive_public};
use crate::proto::{IDLE_TIMEOUT_SECS, KEY_LEN, MALFORMED, PURGE_INTERVAL_SECS};
use crate::ratelimit::RateLimiter;
use crate::session::{ServerConfig, Session};
use crate::store::Store;
pub struct ServeArgs {
pub key_path: String,
pub db_path: String,
pub host: String,
pub port: u16,
pub max_envelope: usize,
pub quota: i64,
pub retention_days: i64,
pub invite_token: Option<String>,
pub rate_connections: u32,
pub rate_sends: u32,
}
pub fn run(args: ServeArgs) -> anyhow::Result<()> {
let static_key =
std::fs::read(&args.key_path).map_err(|e| anyhow::anyhow!("cannot read server key: {e}"))?;
anyhow::ensure!(
static_key.len() == KEY_LEN,
"server key must be {KEY_LEN} raw bytes, got {}",
static_key.len()
);
let config = Arc::new(ServerConfig {
max_envelope: args.max_envelope,
quota: args.quota,
invite_token: args.invite_token.map(String::into_bytes),
conn_limiter: RateLimiter::new(args.rate_connections),
send_limiter: RateLimiter::new(args.rate_sends),
});
let retention_secs = args.retention_days * 86400;
let purge_db_path = args.db_path.clone();
std::thread::spawn(move || purge_loop(purge_db_path, retention_secs));
let listener = std::net::TcpListener::bind((args.host.as_str(), args.port))?;
log::info!("listening on {}:{}", args.host, args.port);
let key_array: [u8; KEY_LEN] = static_key.clone().try_into().unwrap();
let public = derive_public(&key_array);
log::info!("server public key: {}", b32(&public));
for incoming in listener.incoming() {
let stream = match incoming {
Ok(s) => s,
Err(e) => {
log::warn!("accept error: {e}");
continue;
}
};
let peer_ip = stream
.peer_addr()
.map(|a| a.ip().to_string())
.unwrap_or_else(|_| "unknown".to_string());
if !config.conn_limiter.allow(&peer_ip) {
log::warn!("rate limited {peer_ip}");
continue;
}
let config = Arc::clone(&config);
let static_key = static_key.clone();
let db_path = args.db_path.clone();
std::thread::spawn(move || {
if let Err(e) = handle_connection(stream, &config, &static_key, &db_path, &peer_ip) {
log::info!("connection error from {peer_ip}: {e}");
}
});
}
Ok(())
}
fn handle_connection(
mut stream: TcpStream,
config: &ServerConfig,
static_key: &[u8],
db_path: &str,
peer_ip: &str,
) -> anyhow::Result<()> {
stream.set_read_timeout(Some(Duration::from_secs(IDLE_TIMEOUT_SECS)))?;
let (transport, handshake_hash) = match handshake(&mut stream, static_key) {
Ok(v) => v,
Err(e) => {
log::info!("handshake failed from {peer_ip}: {e}");
return Ok(());
}
};
let store = Store::open(db_path)?;
let mut session = Session::new(config, store, peer_ip.to_string(), handshake_hash);
let mut channel = Channel::new(stream, transport);
loop {
let (op, body) = match channel.read_frame() {
Ok(v) => v,
Err(e) => {
if is_eof_like(&e) {
return Ok(());
}
log::info!("bad frame from {peer_ip}: {e}");
let _ = channel.write_frame(0, &[MALFORMED]);
return Ok(());
}
};
let (status, payload) = session.dispatch(op, &body);
let mut response = Vec::with_capacity(1 + payload.len());
response.push(status);
response.extend_from_slice(&payload);
if channel.write_frame(op, &response).is_err() {
return Ok(());
}
}
}
fn is_eof_like(e: &anyhow::Error) -> bool {
if let Some(io_err) = e.downcast_ref::<std::io::Error>() {
return matches!(
io_err.kind(),
std::io::ErrorKind::UnexpectedEof
| std::io::ErrorKind::ConnectionReset
| std::io::ErrorKind::BrokenPipe
| std::io::ErrorKind::TimedOut
| std::io::ErrorKind::WouldBlock
);
}
false
}
fn purge_loop(db_path: String, retention_secs: i64) {
let store = match Store::open(&db_path) {
Ok(s) => s,
Err(e) => {
log::error!("purge thread failed to open store: {e}");
return;
}
};
loop {
std::thread::sleep(Duration::from_secs(PURGE_INTERVAL_SECS));
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs() as i64;
match store.purge(now - retention_secs) {
Ok(0) => {}
Ok(n) => log::info!("expired {n} message(s)"),
Err(e) => log::error!("purge failed: {e}"),
}
}
}

264
src/session.rs Normal file
View file

@ -0,0 +1,264 @@
//! Per-connection dispatch and the six wire operations.
use crate::crypto::{message_id, verify};
use crate::proto::{
valid_username, ProtocolError, Reader, AUTH_FAILED, AUTH_REQUIRED, BAD_VERSION, CERT_LEN,
ENVELOPE_MAGIC, ENVELOPE_MIN, ENVELOPE_VERSION, FETCH_BUDGET, ID_LEN, KEY_LEN, LABEL_AUTH,
LABEL_ROTATE, MALFORMED, MAX_CHAIN, NOT_PERMITTED, OK, OP_AUTH, OP_DELETE, OP_FETCH,
OP_REGISTER, OP_RESOLVE, OP_SEND, QUOTA_EXCEEDED, RATE_LIMITED, TOO_LARGE, UNKNOWN_USER,
};
use crate::ratelimit::RateLimiter;
use crate::store::Store;
pub struct ServerConfig {
pub max_envelope: usize,
pub quota: i64,
pub invite_token: Option<Vec<u8>>,
pub conn_limiter: RateLimiter,
pub send_limiter: RateLimiter,
}
/// A parse failure (-> MALFORMED) or a storage failure (-> INTERNAL_ERROR).
pub enum HandlerError {
Protocol(ProtocolError),
Store(rusqlite::Error),
}
impl From<ProtocolError> for HandlerError {
fn from(e: ProtocolError) -> Self {
HandlerError::Protocol(e)
}
}
impl From<rusqlite::Error> for HandlerError {
fn from(e: rusqlite::Error) -> Self {
HandlerError::Store(e)
}
}
type OpResult = Result<(u8, Vec<u8>), HandlerError>;
pub struct Session<'a> {
config: &'a ServerConfig,
store: Store,
peer_ip: String,
handshake_hash: Vec<u8>,
username: Option<String>,
}
impl<'a> Session<'a> {
pub fn new(config: &'a ServerConfig, store: Store, peer_ip: String, handshake_hash: Vec<u8>) -> Self {
Session {
config,
store,
peer_ip,
handshake_hash,
username: None,
}
}
/// Dispatches one frame, always producing a status to send back, never
/// panicking or propagating errors to the caller: a bad frame or a
/// storage error both become a response, and the caller decides
/// separately whether to keep the connection open.
pub fn dispatch(&mut self, op: u8, body: &[u8]) -> (u8, Vec<u8>) {
if matches!(op, OP_FETCH | OP_DELETE) && self.username.is_none() {
return (AUTH_REQUIRED, Vec::new());
}
let mut r = Reader::new(body);
let result = match op {
OP_AUTH => self.op_auth(&mut r),
OP_RESOLVE => self.op_resolve(&mut r),
OP_SEND => self.op_send(&mut r),
OP_FETCH => self.op_fetch(&mut r),
OP_DELETE => self.op_delete(&mut r),
OP_REGISTER => self.op_register(&mut r),
_ => return (MALFORMED, Vec::new()),
};
match result {
Ok(response) => response,
Err(HandlerError::Protocol(e)) => {
log::info!("bad body from {}: {}", self.peer_ip, e);
(MALFORMED, Vec::new())
}
Err(HandlerError::Store(e)) => {
log::error!("storage error from {}: {}", self.peer_ip, e);
(crate::proto::INTERNAL_ERROR, Vec::new())
}
}
}
fn read_str(r: &mut Reader) -> Result<String, ProtocolError> {
let len = r.u8()? as usize;
let bytes = r.take(len)?;
std::str::from_utf8(bytes)
.map(str::to_string)
.map_err(|_| ProtocolError::new("invalid utf-8"))
}
fn op_auth(&mut self, r: &mut Reader) -> OpResult {
let username = Self::read_str(r)?;
let identity = r.take(KEY_LEN)?.to_vec();
let signature = r.take(64)?.to_vec();
r.done()?;
let 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.as_deref() != Some(identity.as_slice()) {
return Ok((AUTH_FAILED, Vec::new()));
}
let mut msg = LABEL_AUTH.to_vec();
msg.extend_from_slice(&self.handshake_hash);
if !verify(&identity, &signature, &msg) {
return Ok((AUTH_FAILED, Vec::new()));
}
self.username = Some(username);
Ok((OK, Vec::new()))
}
fn op_resolve(&mut self, r: &mut Reader) -> OpResult {
let username = Self::read_str(r)?;
r.done()?;
let identity = match self.store.identity_of(&username)? {
Some(id) => id,
None => return Ok((UNKNOWN_USER, Vec::new())),
};
let chain = self.store.chain(&username)?;
let mut out = identity;
out.push(chain.len() as u8);
for cert in chain {
out.extend_from_slice(&cert);
}
Ok((OK, out))
}
fn op_send(&mut self, r: &mut Reader) -> OpResult {
let envelope = r.rest().to_vec();
if !self.config.send_limiter.allow(&self.peer_ip) {
return Ok((RATE_LIMITED, Vec::new()));
}
if envelope.len() > self.config.max_envelope {
return Ok((TOO_LARGE, Vec::new()));
}
if envelope.len() < ENVELOPE_MIN || &envelope[..4] != ENVELOPE_MAGIC {
return Ok((MALFORMED, Vec::new()));
}
if envelope[4] != ENVELOPE_VERSION {
return Ok((BAD_VERSION, Vec::new()));
}
let recipient = &envelope[5..37];
let username = match self.store.username_for_key(recipient)? {
Some(u) => u,
None => return Ok((UNKNOWN_USER, Vec::new())),
};
let keys = self.store.keys_of(&username)?;
let used = self.store.mailbox_bytes(&keys)?;
if used + envelope.len() as i64 > self.config.quota {
return Ok((QUOTA_EXCEEDED, Vec::new()));
}
// The ciphertext is never inspected; the server cannot read it.
let mid = message_id(&envelope);
self.store.store_message(&mid, recipient, &envelope)?;
Ok((OK, mid.to_vec()))
}
fn op_fetch(&mut self, r: &mut Reader) -> OpResult {
r.done()?;
let username = self.username.as_ref().expect("AUTH_REQUIRED gate above");
let keys = self.store.keys_of(username)?;
let records = self.store.pending(&keys, FETCH_BUDGET)?;
let mut out = Vec::new();
out.extend_from_slice(&(records.len() as u16).to_be_bytes());
for (mid, received_at, envelope) in records {
out.extend_from_slice(&mid);
out.extend_from_slice(&received_at.to_be_bytes());
out.extend_from_slice(&(envelope.len() as u32).to_be_bytes());
out.extend_from_slice(&envelope);
}
Ok((OK, out))
}
fn op_delete(&mut self, r: &mut Reader) -> OpResult {
let count = r.u16()? as usize;
let mut ids = Vec::with_capacity(count);
for _ in 0..count {
ids.push(r.take(ID_LEN)?.to_vec());
}
r.done()?;
let username = self.username.as_ref().expect("AUTH_REQUIRED gate above");
if ids.is_empty() {
return Ok((OK, 0u16.to_be_bytes().to_vec()));
}
// Scoped to the caller's own keys, so ids cannot be used to probe or
// delete another mailbox.
let keys = self.store.keys_of(username)?;
let removed = self.store.delete(&keys, &ids)?;
Ok((OK, (removed as u16).to_be_bytes().to_vec()))
}
fn op_register(&mut self, r: &mut Reader) -> OpResult {
let username = Self::read_str(r)?;
let identity = r.take(KEY_LEN)?.to_vec();
let token_len = r.u8()? as usize;
let token = r.take(token_len)?.to_vec();
let cert_len = r.u8()? as usize;
let cert = r.take(cert_len)?.to_vec();
r.done()?;
if !valid_username(&username) {
return Ok((MALFORMED, Vec::new()));
}
// identity is exactly KEY_LEN bytes by construction (Reader::take
// enforces it); no separate curve-point validity check is needed.
if let Some(expected) = &self.config.invite_token {
if &token != expected {
return Ok((NOT_PERMITTED, Vec::new()));
}
}
if cert.is_empty() {
if !self.store.register(&username, &identity)? {
return Ok((NOT_PERMITTED, Vec::new()));
}
return Ok((OK, Vec::new()));
}
if cert.len() != CERT_LEN {
return Ok((MALFORMED, Vec::new()));
}
let old_pub = &cert[..32];
let new_pub = &cert[32..64];
let when = &cert[64..72];
let signature = &cert[72..];
if new_pub != identity.as_slice() {
return Ok((MALFORMED, Vec::new()));
}
let bound = match self.store.identity_of(&username)? {
Some(b) => b,
None => return Ok((UNKNOWN_USER, Vec::new())),
};
// Only the currently bound key may hand the username on.
if bound != old_pub {
return Ok((NOT_PERMITTED, Vec::new()));
}
let mut msg = LABEL_ROTATE.to_vec();
msg.extend_from_slice(old_pub);
msg.extend_from_slice(new_pub);
msg.extend_from_slice(when);
if !verify(old_pub, signature, &msg) {
return Ok((AUTH_FAILED, Vec::new()));
}
let chain = self.store.chain(&username)?;
if chain.len() >= MAX_CHAIN {
return Ok((NOT_PERMITTED, Vec::new()));
}
if !self.store.rotate(&username, new_pub, &cert, chain.len())? {
return Ok((NOT_PERMITTED, Vec::new()));
}
Ok((OK, Vec::new()))
}
}

243
src/store.rs Normal file
View file

@ -0,0 +1,243 @@
//! SQLite-backed mailbox storage.
//!
//! Each connection thread opens its own `Store` (own `rusqlite::Connection`),
//! since SQLite connections aren't meant to be shared across threads.
use std::time::Duration;
use rusqlite::{params_from_iter, Connection, OptionalExtension};
const SCHEMA: &str = "
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.
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);
";
pub struct Store {
conn: Connection,
}
fn placeholders(n: usize) -> String {
vec!["?"; n].join(",")
}
impl Store {
pub fn open(path: &str) -> rusqlite::Result<Self> {
let conn = Connection::open(path)?;
conn.busy_timeout(Duration::from_secs(10))?;
conn.pragma_update(None, "journal_mode", "WAL")?;
conn.execute_batch(SCHEMA)?;
Ok(Store { conn })
}
pub fn identity_of(&self, username: &str) -> rusqlite::Result<Option<Vec<u8>>> {
self.conn
.query_row(
"SELECT identity FROM users WHERE username = ?1",
[username],
|row| row.get(0),
)
.optional()
}
pub fn chain(&self, username: &str) -> rusqlite::Result<Vec<Vec<u8>>> {
let mut stmt = self
.conn
.prepare("SELECT cert FROM rotations WHERE username = ?1 ORDER BY seq")?;
let rows = stmt.query_map([username], |row| row.get(0))?;
rows.collect()
}
pub fn keys_of(&self, username: &str) -> rusqlite::Result<Vec<Vec<u8>>> {
let mut stmt = self
.conn
.prepare("SELECT identity FROM keys WHERE username = ?1")?;
let rows = stmt.query_map([username], |row| row.get(0))?;
rows.collect()
}
pub fn username_for_key(&self, identity: &[u8]) -> rusqlite::Result<Option<String>> {
self.conn
.query_row(
"SELECT username FROM keys WHERE identity = ?1",
[identity],
|row| row.get(0),
)
.optional()
}
/// False on conflict: username taken, or this key is already bound elsewhere.
pub fn register(&self, username: &str, identity: &[u8]) -> rusqlite::Result<bool> {
let tx = self.conn.unchecked_transaction()?;
let result = (|| -> rusqlite::Result<()> {
tx.execute(
"INSERT INTO users (username, identity) VALUES (?1, ?2)",
(username, identity),
)?;
tx.execute(
"INSERT INTO keys (identity, username) VALUES (?1, ?2)",
(identity, username),
)?;
Ok(())
})();
match result {
Ok(()) => {
tx.commit()?;
Ok(true)
}
Err(rusqlite::Error::SqliteFailure(e, _))
if e.code == rusqlite::ErrorCode::ConstraintViolation =>
{
Ok(false)
}
Err(e) => Err(e),
}
}
pub fn rotate(
&self,
username: &str,
new_key: &[u8],
cert: &[u8],
seq: usize,
) -> rusqlite::Result<bool> {
let tx = self.conn.unchecked_transaction()?;
let result = (|| -> rusqlite::Result<()> {
tx.execute(
"UPDATE users SET identity = ?1 WHERE username = ?2",
(new_key, username),
)?;
tx.execute(
"INSERT INTO keys (identity, username) VALUES (?1, ?2)",
(new_key, username),
)?;
tx.execute(
"INSERT INTO rotations (username, seq, cert) VALUES (?1, ?2, ?3)",
(username, seq as i64, cert),
)?;
Ok(())
})();
match result {
Ok(()) => {
tx.commit()?;
Ok(true)
}
Err(rusqlite::Error::SqliteFailure(e, _))
if e.code == rusqlite::ErrorCode::ConstraintViolation =>
{
Ok(false)
}
Err(e) => Err(e),
}
}
pub fn mailbox_bytes(&self, keys: &[Vec<u8>]) -> rusqlite::Result<i64> {
if keys.is_empty() {
return Ok(0);
}
let sql = format!(
"SELECT COALESCE(SUM(LENGTH(envelope)), 0) FROM messages WHERE recipient IN ({})",
placeholders(keys.len())
);
self.conn
.query_row(&sql, params_from_iter(keys.iter()), |row| row.get(0))
}
pub fn store_message(
&self,
mid: &[u8],
recipient: &[u8],
envelope: &[u8],
) -> rusqlite::Result<()> {
self.conn.execute(
"INSERT OR IGNORE INTO messages (id, recipient, received_at, envelope) \
VALUES (?1, ?2, ?3, ?4)",
(
mid,
recipient,
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs() as i64,
envelope,
),
)?;
Ok(())
}
/// Always returns at least one message, even if it alone exceeds `budget`,
/// so an oversized envelope cannot wedge a mailbox shut.
pub fn pending(
&self,
keys: &[Vec<u8>],
budget: usize,
) -> rusqlite::Result<Vec<(Vec<u8>, i64, Vec<u8>)>> {
if keys.is_empty() {
return Ok(Vec::new());
}
let sql = format!(
"SELECT id, received_at, envelope FROM messages \
WHERE recipient IN ({}) ORDER BY received_at, id",
placeholders(keys.len())
);
let mut stmt = self.conn.prepare(&sql)?;
let rows = stmt.query_map(params_from_iter(keys.iter()), |row| {
Ok((
row.get::<_, Vec<u8>>(0)?,
row.get::<_, i64>(1)?,
row.get::<_, Vec<u8>>(2)?,
))
})?;
let mut out = Vec::new();
let mut used = 0usize;
for row in rows {
let (mid, received_at, envelope) = row?;
if !out.is_empty() && used + envelope.len() > budget {
break;
}
used += envelope.len();
out.push((mid, received_at, envelope));
}
Ok(out)
}
pub fn delete(&self, keys: &[Vec<u8>], ids: &[Vec<u8>]) -> rusqlite::Result<usize> {
if keys.is_empty() || ids.is_empty() {
return Ok(0);
}
let sql = format!(
"DELETE FROM messages WHERE id IN ({}) AND recipient IN ({})",
placeholders(ids.len()),
placeholders(keys.len())
);
let params: Vec<&Vec<u8>> = ids.iter().chain(keys.iter()).collect();
self.conn.execute(&sql, params_from_iter(params))
}
pub fn purge(&self, older_than: i64) -> rusqlite::Result<usize> {
self.conn
.execute("DELETE FROM messages WHERE received_at < ?1", [older_than])
}
}