bunshin/src/server.rs

354 lines
11 KiB
Rust

//! TCP accept loop, per-connection handling, and the background purge loop.
use std::net::{TcpListener, TcpStream};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use crate::bind::TransportBindValues;
use crate::channel::{handshake, Channel};
use crate::config::DomainConfig;
use crate::crypto::{b32, derive_public};
use crate::proto::{FETCH_BUDGET, HEADER_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 requests_quota: i64,
pub retention_days: i64,
pub requests_retention_days: i64,
pub max_tokens: u16,
pub invite_token: Option<String>,
pub rate_connections: u32,
pub rate_sends: u32,
pub rate_tokens: u32,
pub max_connections: usize,
}
/// The single-domain CLI path: one anonymous domain.
pub fn run(args: ServeArgs) -> anyhow::Result<()> {
serve_domains(vec![DomainConfig {
name: "default".to_string(),
key_path: args.key_path,
db_path: args.db_path,
host: args.host,
port: args.port,
max_envelope: args.max_envelope,
quota: args.quota,
requests_quota: args.requests_quota,
retention_days: args.retention_days,
requests_retention_days: args.requests_retention_days,
max_tokens: args.max_tokens,
invite_token: args.invite_token.map(String::into_bytes),
rate_connections: args.rate_connections,
rate_sends: args.rate_sends,
rate_tokens: args.rate_tokens,
max_connections: args.max_connections,
}])
}
/// Serves every domain: one listener, key, mailbox database, config and
/// purge loop per domain, nothing shared between them.
pub fn serve_domains(domains: Vec<DomainConfig>) -> anyhow::Result<()> {
// Read every key and bind every listener first, so a missing key or a
// taken port fails the whole process before any domain serves.
let mut bound = Vec::with_capacity(domains.len());
for domain in &domains {
let static_key = read_static_key(&domain.key_path)?;
let listener = TcpListener::bind((domain.host.as_str(), domain.port))?;
let config = Arc::new(ServerConfig {
max_envelope: domain.max_envelope,
fetch_budget: FETCH_BUDGET,
main_quota: domain.quota,
requests_quota: domain.requests_quota,
max_tokens: domain.max_tokens,
invite_token: domain.invite_token.clone(),
conn_limiter: RateLimiter::new(domain.rate_connections),
send_limiter: RateLimiter::new(domain.rate_sends),
token_limiter: RateLimiter::new(domain.rate_tokens),
});
bound.push((domain, static_key, listener, config));
}
for (domain, static_key, listener, config) in bound {
log::info!(
"domain {}: listening on {}:{}",
domain.name,
domain.host,
domain.port
);
log::info!(
"domain {}: server public key: {}",
domain.name,
b32(&derive_public(&static_key))
);
let purge_db_path = domain.db_path.clone();
let main_retention_secs = domain.retention_days * 86400;
let requests_retention_secs = domain.requests_retention_days * 86400;
std::thread::spawn(move || {
purge_loop(purge_db_path, main_retention_secs, requests_retention_secs)
});
let gate = Arc::new(ConnGate::new(domain.max_connections));
let db_path = domain.db_path.clone();
std::thread::spawn(move || accept_loop(listener, config, static_key, db_path, gate));
}
// Each domain's accept loop runs in its own thread; nothing fails here.
loop {
std::thread::sleep(Duration::from_secs(3600));
}
}
fn read_static_key(path: &str) -> anyhow::Result<[u8; KEY_LEN]> {
let key =
std::fs::read(path).map_err(|e| anyhow::anyhow!("cannot read server key {path}: {e}"))?;
anyhow::ensure!(
key.len() == KEY_LEN,
"server key {path} must be {KEY_LEN} raw bytes, got {}",
key.len()
);
Ok(key.try_into().unwrap())
}
fn accept_loop(
listener: TcpListener,
config: Arc<ServerConfig>,
static_key: [u8; KEY_LEN],
db_path: String,
gate: Arc<ConnGate>,
) {
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 permit = match gate.try_acquire() {
Some(p) => p,
None => {
log::warn!("connection limit reached, refusing {peer_ip}");
continue;
}
};
let config = Arc::clone(&config);
let db_path = db_path.clone();
std::thread::spawn(move || {
// The permit is dropped with the connection, freeing its slot.
let _permit = permit;
if let Err(e) = handle_connection(stream, &config, &static_key, &db_path, &peer_ip) {
log::info!("connection error from {peer_ip}: {e}");
}
});
}
}
/// Cap on concurrently served TCP connections, the mirror of the RNS
/// carrier's link cap (SPEC.md sec 13.8): without one, a connection flood
/// would exhaust threads one unbounded spawn at a time. `max == 0` disables
/// the cap; past it, connections are refused, never queued.
struct ConnGate {
max: usize,
active: Mutex<usize>,
}
impl ConnGate {
fn new(max: usize) -> Self {
ConnGate {
max,
active: Mutex::new(0),
}
}
fn try_acquire(self: &Arc<Self>) -> Option<ConnPermit> {
if self.max == 0 {
return Some(ConnPermit {
gate: Arc::clone(self),
});
}
let mut active = self.active.lock().unwrap();
if *active >= self.max {
return None;
}
*active += 1;
Some(ConnPermit {
gate: Arc::clone(self),
})
}
}
/// Dropping releases the slot, so the count stays accurate whatever return
/// path or panic closes the connection.
struct ConnPermit {
gate: Arc<ConnGate>,
}
impl Drop for ConnPermit {
fn drop(&mut self) {
if self.gate.max == 0 {
return;
}
*self.gate.active.lock().unwrap() -= 1;
}
}
/// Also the hostile harness's entry point: it drives real connections
/// through the same accept/handshake/session path the TCP carrier serves.
pub(crate) fn handle_connection(
mut stream: TcpStream,
config: &ServerConfig,
static_key: &[u8; KEY_LEN],
db_path: &str,
peer_ip: &str,
) -> anyhow::Result<()> {
// The handshake runs under the header budget; once it completes,
// Channel::read_noise picks between the header and idle deadlines per
// frame state.
stream.set_read_timeout(Some(Duration::from_secs(HEADER_TIMEOUT_SECS)))?;
let handshake_started = std::time::Instant::now();
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(());
}
};
log::debug!(
"handshake from {peer_ip} in {:?}",
handshake_started.elapsed()
);
// Drop logs the summary once, whatever return path closes the session.
let summary = SessionSummary {
peer_ip,
started: std::time::Instant::now(),
ops: std::cell::Cell::new(0),
};
let store = Store::open(db_path)?;
let server_static = derive_public(static_key);
let bind = TransportBindValues::tcp(&handshake_hash, &server_static)?;
let mut session = Session::new(config, store, peer_ip.to_string(), bind);
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(());
}
};
summary.ops.set(summary.ops.get() + 1);
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(());
}
}
}
struct SessionSummary<'a> {
peer_ip: &'a str,
started: std::time::Instant,
ops: std::cell::Cell<u64>,
}
impl Drop for SessionSummary<'_> {
fn drop(&mut self) {
log::debug!(
"session {} closed: {} op(s) in {:?}",
self.peer_ip,
self.ops.get(),
self.started.elapsed()
);
}
}
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, main_retention_secs: i64, requests_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 - main_retention_secs, now - requests_retention_secs) {
Ok(0) => {}
Ok(n) => log::info!("expired {n} message(s)"),
Err(e) => log::error!("purge failed: {e}"),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn gate_refuses_past_cap_and_frees_on_drop() {
let gate = Arc::new(ConnGate::new(2));
let a = gate.try_acquire().unwrap();
let b = gate.try_acquire().unwrap();
assert!(gate.try_acquire().is_none());
drop(b);
assert!(gate.try_acquire().is_some());
drop(a);
}
#[test]
fn zero_cap_means_unlimited() {
let gate = Arc::new(ConnGate::new(0));
for _ in 0..100 {
assert!(gate.try_acquire().is_some());
}
}
}