pub mod auth;
mod conn;
mod convert;
#[cfg(feature = "metrics")]
mod metrics;
mod stats;
#[cfg(feature = "tls")]
pub mod tls;
use std::collections::HashMap;
use std::io;
use std::net::{Shutdown, SocketAddr, TcpListener, TcpStream, ToSocketAddrs};
use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::{Duration, Instant};
use socket2::{SockRef, TcpKeepalive};
use tephra::writer::WriteHandle;
use tephra_proto::DEFAULT_MAX_FRAME_LEN;
pub use convert::ConvertError;
#[derive(Clone, Copy, Debug)]
pub struct ServerConfig {
pub max_frame_len: u32,
pub read_batch_events: usize,
pub read_batch_bytes: usize,
pub subscribe_wait_tick: Duration,
pub max_inflight_requests_per_conn: usize,
pub max_concurrent_subscriptions: usize,
pub read_worker_threads: usize,
pub frame_queue_depth: usize,
pub keepalive_idle: Duration,
pub keepalive_interval: Duration,
pub max_connections: usize,
pub incomplete_frame_timeout: Duration,
pub handshake_timeout: Duration,
pub idle_timeout: Duration,
}
impl Default for ServerConfig {
fn default() -> Self {
ServerConfig {
max_frame_len: DEFAULT_MAX_FRAME_LEN,
read_batch_events: 1024,
read_batch_bytes: 512 * 1024,
subscribe_wait_tick: Duration::from_millis(250),
max_inflight_requests_per_conn: 256,
max_concurrent_subscriptions: 64,
read_worker_threads: 0,
frame_queue_depth: 256,
keepalive_idle: Duration::from_secs(60),
keepalive_interval: Duration::from_secs(15),
max_connections: 1024,
incomplete_frame_timeout: Duration::from_secs(30),
handshake_timeout: Duration::ZERO,
idle_timeout: Duration::ZERO,
}
}
}
#[derive(Clone, Default)]
struct Connections {
inner: Arc<Mutex<ConnectionsInner>>,
}
#[derive(Default)]
struct ConnectionsInner {
shutting_down: bool,
streams: HashMap<u64, TcpStream>,
}
impl Connections {
fn register(&self, id: u64, stream: TcpStream) -> bool {
let mut inner = self.inner.lock().unwrap();
if inner.shutting_down {
let _ = stream.shutdown(Shutdown::Both);
return false;
}
inner.streams.insert(id, stream);
true
}
fn remove(&self, id: u64) {
self.inner.lock().unwrap().streams.remove(&id);
}
fn shutdown_all(&self) {
let mut inner = self.inner.lock().unwrap();
inner.shutting_down = true;
for stream in inner.streams.values() {
let _ = stream.shutdown(Shutdown::Both);
}
}
}
pub(crate) struct SharedStats {
data_dir: Option<PathBuf>,
start_time: Instant,
active_connections: AtomicU64,
active_subscriptions: AtomicU64,
connections_refused: AtomicU64,
connections_reaped: AtomicU64,
max_connections: u64,
}
impl SharedStats {
fn new(data_dir: Option<PathBuf>, max_connections: u64) -> SharedStats {
SharedStats {
data_dir,
start_time: Instant::now(),
active_connections: AtomicU64::new(0),
active_subscriptions: AtomicU64::new(0),
connections_refused: AtomicU64::new(0),
connections_reaped: AtomicU64::new(0),
max_connections,
}
}
}
struct ConnPermit(Arc<SharedStats>);
impl ConnPermit {
fn acquire(stats: &Arc<SharedStats>) -> Option<ConnPermit> {
let max = stats.max_connections;
if max != 0 && stats.active_connections.load(Ordering::Relaxed) >= max {
return None;
}
stats.active_connections.fetch_add(1, Ordering::Relaxed);
Some(ConnPermit(Arc::clone(stats)))
}
}
impl Drop for ConnPermit {
fn drop(&mut self) {
self.0.active_connections.fetch_sub(1, Ordering::Relaxed);
}
}
pub struct Server {
listener: TcpListener,
handle: WriteHandle,
config: ServerConfig,
local_addr: SocketAddr,
running: Arc<AtomicBool>,
connections: Connections,
data_dir: Option<PathBuf>,
#[cfg(feature = "metrics")]
metrics_listener: Option<TcpListener>,
#[cfg(feature = "metrics")]
metrics_addr: Option<SocketAddr>,
#[cfg(feature = "tls")]
tls: Option<Arc<rustls::ServerConfig>>,
auth: Option<Arc<auth::AuthConfig>>,
}
impl Server {
pub fn bind(
addr: impl ToSocketAddrs,
handle: WriteHandle,
config: ServerConfig,
) -> io::Result<Server> {
let listener = TcpListener::bind(addr)?;
let local_addr = listener.local_addr()?;
Ok(Server {
listener,
handle,
config,
local_addr,
running: Arc::new(AtomicBool::new(true)),
connections: Connections::default(),
data_dir: None,
#[cfg(feature = "metrics")]
metrics_listener: None,
#[cfg(feature = "metrics")]
metrics_addr: None,
#[cfg(feature = "tls")]
tls: None,
auth: None,
})
}
#[cfg(feature = "tls")]
pub fn with_tls(mut self, tls_config: Arc<rustls::ServerConfig>) -> Server {
self.tls = Some(tls_config);
self
}
pub fn with_auth(mut self, auth: Arc<auth::AuthConfig>) -> Server {
self.auth = Some(auth);
self
}
pub fn with_data_dir(mut self, data_dir: impl Into<PathBuf>) -> Server {
self.data_dir = Some(data_dir.into());
self
}
#[cfg(feature = "metrics")]
pub fn with_metrics_addr(mut self, addr: impl ToSocketAddrs) -> io::Result<Server> {
let listener = TcpListener::bind(addr)?;
self.metrics_addr = Some(listener.local_addr()?);
self.metrics_listener = Some(listener);
Ok(self)
}
#[cfg(feature = "metrics")]
pub fn metrics_local_addr(&self) -> Option<SocketAddr> {
self.metrics_addr
}
pub fn local_addr(&self) -> SocketAddr {
self.local_addr
}
pub fn shutdown_handle(&self) -> ShutdownHandle {
ShutdownHandle {
running: Arc::clone(&self.running),
local_addr: self.local_addr,
connections: self.connections.clone(),
}
}
pub fn run(self) -> io::Result<()> {
tracing::info!(addr = %self.local_addr, "tephra server listening");
let next_id = AtomicU64::new(0);
let mut threads = Vec::new();
let stats = Arc::new(SharedStats::new(
self.data_dir.clone(),
self.config.max_connections as u64,
));
#[cfg(feature = "tls")]
let tls_acceptor = self.tls.clone();
let auth = self.auth.clone();
#[cfg(feature = "metrics")]
let metrics_thread = match self.metrics_listener {
Some(listener) => {
if let Some(addr) = self.metrics_addr {
tracing::info!(%addr, "metrics endpoint listening");
}
let handle = self.handle.clone();
let stats = Arc::clone(&stats);
let running = Arc::clone(&self.running);
Some(
thread::Builder::new()
.name("tephra-metrics".to_string())
.spawn(move || metrics::serve(listener, handle, stats, running))
.expect("spawn metrics thread"),
)
}
None => None,
};
let read_workers = if self.config.read_worker_threads != 0 {
self.config.read_worker_threads
} else {
thread::available_parallelism()
.map(|n| n.get())
.unwrap_or(1)
};
let read_pool = conn::ReadPool::new(read_workers);
tracing::info!(read_workers, "read worker pool started");
let mut at_cap = false;
for stream in self.listener.incoming() {
if !self.running.load(Ordering::Acquire) {
break;
}
let stream = match stream {
Ok(stream) => stream,
Err(err) => {
tracing::warn!(%err, "accept failed");
continue;
}
};
let permit = match ConnPermit::acquire(&stats) {
Some(permit) => {
at_cap = false;
permit
}
None => {
let refused = stats.connections_refused.fetch_add(1, Ordering::Relaxed) + 1;
if !at_cap {
at_cap = true;
tracing::warn!(
max = self.config.max_connections,
refused_total = refused,
"at max_connections; refusing new connections until one frees"
);
}
let _ = stream.shutdown(Shutdown::Both);
continue;
}
};
threads.retain(|thread: &thread::JoinHandle<()>| !thread.is_finished());
let keepalive = TcpKeepalive::new()
.with_time(self.config.keepalive_idle)
.with_interval(self.config.keepalive_interval);
if let Err(err) = SockRef::from(&stream).set_tcp_keepalive(&keepalive) {
tracing::warn!(%err, "failed to set TCP keepalive");
}
let id = next_id.fetch_add(1, Ordering::Relaxed);
let registry_handle = match stream.try_clone() {
Ok(clone) => clone,
Err(err) => {
tracing::warn!(%err, "failed to clone connection for shutdown; dropping it");
continue;
}
};
if !self.connections.register(id, registry_handle) {
let _ = stream.shutdown(Shutdown::Both);
break;
}
let handle = self.handle.clone();
let config = self.config;
let connections = self.connections.clone();
let running = Arc::clone(&self.running);
let read_pool = read_pool.sender();
let auth = auth.clone();
#[cfg(feature = "tls")]
let tls = tls_acceptor.clone();
let thread = thread::Builder::new()
.name("tephra-conn".to_string())
.spawn(move || {
conn::serve_connection(
stream,
handle,
config,
running,
read_pool,
&permit.0,
auth,
#[cfg(feature = "tls")]
tls,
);
drop(permit);
connections.remove(id);
})?;
threads.push(thread);
}
for thread in threads {
let _ = thread.join();
}
read_pool.shutdown();
#[cfg(feature = "metrics")]
if let Some(metrics_thread) = metrics_thread {
let _ = metrics_thread.join();
}
tracing::info!("tephra server stopped");
Ok(())
}
}
#[derive(Clone)]
pub struct ShutdownHandle {
running: Arc<AtomicBool>,
local_addr: SocketAddr,
connections: Connections,
}
impl ShutdownHandle {
pub fn shutdown(&self) {
self.running.store(false, Ordering::Release);
let _ = TcpStream::connect(self.local_addr);
self.connections.shutdown_all();
}
}