mod conn;
mod convert;
use std::collections::HashMap;
use std::io;
use std::net::{Shutdown, SocketAddr, TcpListener, TcpStream, ToSocketAddrs};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::Duration;
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,
}
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),
}
}
}
#[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 struct Server {
listener: TcpListener,
handle: WriteHandle,
config: ServerConfig,
local_addr: SocketAddr,
running: Arc<AtomicBool>,
connections: Connections,
}
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(),
})
}
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 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");
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 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 thread = thread::Builder::new()
.name("tephra-conn".to_string())
.spawn(move || {
conn::serve_connection(stream, handle, config, running, read_pool);
connections.remove(id);
})?;
threads.push(thread);
}
for thread in threads {
let _ = thread.join();
}
read_pool.shutdown();
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();
}
}