use std::collections::HashMap;
use std::io::{BufReader, BufWriter, Write};
use std::net::{Shutdown, SocketAddr, TcpStream};
use std::panic::AssertUnwindSafe;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Condvar, Mutex};
use std::thread;
use flume::{Receiver, Sender};
use tephra::log::set::PositionRange;
use tephra::query::Query;
use tephra::read::WaitOutcome;
use tephra::writer::{AppendError, WriteHandle};
use tephra::{Event, Position};
use tephra_proto::tephra as pb;
use tephra_proto::{FrameError, read_frame, write_frame};
use crate::ServerConfig;
use crate::convert;
type AppendReply = (u64, Result<PositionRange, AppendError>);
pub(crate) fn serve_connection(
stream: TcpStream,
handle: WriteHandle,
config: ServerConfig,
running: Arc<AtomicBool>,
read_pool: Sender<ReadJob>,
) {
let peer = stream.peer_addr().ok();
if let Err(err) = stream.set_nodelay(true) {
tracing::warn!(?peer, %err, "failed to set TCP_NODELAY");
}
let read_half = match stream.try_clone() {
Ok(clone) => clone,
Err(err) => {
tracing::warn!(?peer, %err, "failed to clone connection stream");
return;
}
};
let write_half = match stream.try_clone() {
Ok(clone) => clone,
Err(err) => {
tracing::warn!(?peer, %err, "failed to clone connection stream");
return;
}
};
let alive = Arc::new(AtomicBool::new(true));
let (out_tx, out_rx) = flume::bounded::<pb::Response>(config.frame_queue_depth);
let (reply_tx, reply_rx) = flume::unbounded::<AppendReply>();
let cancels: Arc<Mutex<HashMap<u64, Arc<AtomicBool>>>> = Arc::default();
let inflight = Arc::new(Semaphore::new(config.max_inflight_requests_per_conn));
let subscriptions = Arc::new(Semaphore::new(config.max_concurrent_subscriptions));
let workers = WaitGroup::default();
let writer_thread = {
let alive = Arc::clone(&alive);
let shutdown = stream.try_clone().ok();
let max = config.max_frame_len;
thread::Builder::new()
.name("tephra-conn-writer".to_string())
.spawn(move || writer_loop(write_half, out_rx, shutdown, alive, max, peer))
.expect("spawn connection writer thread")
};
let pump_thread = {
let out_tx = out_tx.clone();
let inflight = Arc::clone(&inflight);
thread::Builder::new()
.name("tephra-conn-pump".to_string())
.spawn(move || pump_loop(reply_rx, out_tx, inflight))
.expect("spawn connection pump thread")
};
let conn = ConnCtx {
handle,
config,
running,
alive: Arc::clone(&alive),
out_tx: out_tx.clone(),
cancels: Arc::clone(&cancels),
workers: workers.clone(),
inflight,
subscriptions,
read_pool,
};
let mut reader = BufReader::new(read_half);
loop {
let request = match read_frame::<pb::Request, _>(&mut reader, config.max_frame_len) {
Ok(Some(request)) => request,
Ok(None) => {
tracing::debug!(?peer, "connection closed by peer at a frame boundary");
break;
}
Err(err) => {
let error = match &err {
FrameError::TooLarge { .. } => Some(convert::too_large(err.to_string())),
FrameError::Parse(_) => Some(convert::bad_request(err.to_string())),
FrameError::Io(_) | FrameError::Serialize(_) => None,
};
if let Some(error) = error {
let _ = out_tx.send(make_response(0, ResponseKind::Error(error)));
}
if alive.load(Ordering::Acquire) {
tracing::warn!(?peer, %err, "closing connection: reader failed");
} else {
tracing::debug!(?peer, %err, "reader ended after the connection was closed");
}
break;
}
};
dispatch(&request, &conn, &reply_tx);
}
alive.store(false, Ordering::Release);
drop(conn);
drop(out_tx);
drop(reply_tx);
workers.wait();
let _ = pump_thread.join();
let _ = writer_thread.join();
let _ = stream.shutdown(Shutdown::Both);
}
#[derive(Clone)]
struct ConnCtx {
handle: WriteHandle,
config: ServerConfig,
running: Arc<AtomicBool>,
alive: Arc<AtomicBool>,
out_tx: Sender<pb::Response>,
cancels: Arc<Mutex<HashMap<u64, Arc<AtomicBool>>>>,
workers: WaitGroup,
inflight: Arc<Semaphore>,
subscriptions: Arc<Semaphore>,
read_pool: Sender<ReadJob>,
}
impl ConnCtx {
fn should_continue(&self, cancel: &AtomicBool) -> bool {
self.running.load(Ordering::Acquire)
&& self.alive.load(Ordering::Acquire)
&& !cancel.load(Ordering::Acquire)
}
}
fn dispatch(request: &pb::Request, conn: &ConnCtx, reply_tx: &Sender<AppendReply>) {
let request_id = request.request_id();
match request.kind() {
pb::request::KindOneof::Append(append) => handle_append(request_id, append, conn, reply_tx),
pb::request::KindOneof::Read(read) => spawn_read(request_id, read, conn),
pb::request::KindOneof::Subscribe(subscribe) => {
spawn_subscribe(request_id, subscribe, conn)
}
pb::request::KindOneof::Cancel(cancel) => {
if let Some(flag) = conn.cancels.lock().unwrap().get(&cancel.target()) {
flag.store(true, Ordering::Release);
}
}
_ => {
let error =
convert::bad_request("request has no append, read, subscribe, or cancel set");
let _ = conn
.out_tx
.send(make_response(request_id, ResponseKind::Error(error)));
}
}
}
fn handle_append(
request_id: u64,
append: pb::AppendRequestView<'_>,
conn: &ConnCtx,
reply_tx: &Sender<AppendReply>,
) {
let events = match convert::events_from_proto(append) {
Ok(events) => events,
Err(err) => {
let _ = conn.out_tx.send(make_response(
request_id,
ResponseKind::Error(convert::bad_request(err)),
));
return;
}
};
let condition = match append.condition_opt() {
Some(condition) => match convert::condition_from_proto(condition) {
Ok(condition) => Some(condition),
Err(err) => {
let _ = conn.out_tx.send(make_response(
request_id,
ResponseKind::Error(convert::bad_request(err)),
));
return;
}
},
None => None,
};
conn.inflight.acquire();
if let Err(err) = conn
.handle
.append_submit(events, condition, request_id, reply_tx.clone())
{
conn.inflight.release();
let _ = conn.out_tx.send(make_response(
request_id,
ResponseKind::Error(convert::append_error_to_proto(&err)),
));
}
}
fn spawn_read(request_id: u64, read: pb::ReadRequestView<'_>, conn: &ConnCtx) {
let query = match convert::query_from_proto(read.query()) {
Ok(query) => query,
Err(err) => {
let _ = conn.out_tx.send(make_response(
request_id,
ResponseKind::Error(convert::bad_request(err)),
));
return;
}
};
let after = Position::new(read.after());
let limit = read.limit_opt();
let cancel = register_cancel(conn, request_id);
conn.inflight.acquire();
conn.workers.add();
let cleanup = WorkerCleanup {
cancels: Arc::clone(&conn.cancels),
sem: Some(Arc::clone(&conn.inflight)),
workers: conn.workers.clone(),
request_id,
};
let job = ReadJob {
request_id,
query,
after,
limit,
conn: conn.clone(),
cancel,
cleanup,
};
if conn.read_pool.send(job).is_err() {
tracing::debug!(request_id, "read pool closed; dropping read");
}
}
fn spawn_subscribe(request_id: u64, subscribe: pb::SubscribeRequestView<'_>, conn: &ConnCtx) {
let query = match convert::query_from_proto(subscribe.query()) {
Ok(query) => query,
Err(err) => {
let _ = conn.out_tx.send(make_response(
request_id,
ResponseKind::Error(convert::bad_request(err)),
));
return;
}
};
let after = Position::new(subscribe.after());
if !conn.subscriptions.try_acquire() {
let _ = conn.out_tx.send(make_response(
request_id,
ResponseKind::Error(convert::bad_request(
"too many concurrent subscriptions on this connection",
)),
));
return;
}
let cancel = register_cancel(conn, request_id);
conn.workers.add();
let cleanup = WorkerCleanup {
cancels: Arc::clone(&conn.cancels),
sem: Some(Arc::clone(&conn.subscriptions)),
workers: conn.workers.clone(),
request_id,
};
let conn_owned = conn.clone();
if let Err(err) = thread::Builder::new()
.name("tephra-conn-subscribe".to_string())
.spawn(move || {
let _cleanup = cleanup;
run_subscribe(request_id, query, after, &conn_owned, &cancel);
})
{
tracing::warn!(%err, "failed to spawn subscribe worker");
let _ = conn.out_tx.send(make_response(
request_id,
ResponseKind::Error(convert::bad_request(
"server could not start the subscription",
)),
));
}
}
fn register_cancel(conn: &ConnCtx, request_id: u64) -> Arc<AtomicBool> {
let cancel = Arc::new(AtomicBool::new(false));
conn.cancels
.lock()
.unwrap()
.insert(request_id, Arc::clone(&cancel));
cancel
}
fn run_read(
request_id: u64,
query: Query,
after: Position,
limit: Option<u64>,
conn: &ConnCtx,
cancel: &AtomicBool,
) {
let mut reads = conn.handle.read(query, after, limit);
let watermark = reads.watermark();
let mut batch = pb::ReadEvents::new();
let mut batch_bytes = 0usize;
while let Some(item) = reads.next() {
if !conn.should_continue(cancel) {
return;
}
let sequenced = match item {
Ok(sequenced) => sequenced,
Err(err) => {
let _ = conn.out_tx.send(make_response(
request_id,
ResponseKind::Error(convert::internal_error(err)),
));
return;
}
};
batch_bytes += sequenced.event.as_bytes().len();
batch.events_mut().push(convert::sequenced_to_proto(
sequenced.position,
sequenced.event,
));
if batch.events().len() >= conn.config.read_batch_events
|| batch_bytes >= conn.config.read_batch_bytes
{
let full = std::mem::replace(&mut batch, pb::ReadEvents::new());
if conn
.out_tx
.send(make_response(request_id, ResponseKind::ReadEvents(full)))
.is_err()
{
return;
}
batch_bytes = 0;
}
}
if !batch.events().is_empty()
&& conn
.out_tx
.send(make_response(request_id, ResponseKind::ReadEvents(batch)))
.is_err()
{
return;
}
let mut end = pb::ReadEnd::new();
end.set_watermark(watermark.get());
let _ = conn
.out_tx
.send(make_response(request_id, ResponseKind::ReadEnd(end)));
}
fn run_subscribe(
request_id: u64,
query: Query,
after: Position,
conn: &ConnCtx,
cancel: &AtomicBool,
) {
let mut sub = conn.handle.subscribe(query, after);
let mut announced = false;
loop {
if !conn.should_continue(cancel) {
return;
}
let batch = match sub.poll_batch() {
Ok(batch) => batch,
Err(err) => {
let _ = conn.out_tx.send(make_response(
request_id,
ResponseKind::Error(convert::internal_error(err)),
));
return;
}
};
if batch.is_empty() {
if !announced {
let mut caught_up = pb::SubscribeCaughtUp::new();
caught_up.set_watermark(sub.position().get());
if conn
.out_tx
.send(make_response(request_id, ResponseKind::CaughtUp(caught_up)))
.is_err()
{
return;
}
announced = true;
}
match sub.wait_timeout(conn.config.subscribe_wait_tick) {
WaitOutcome::Advanced | WaitOutcome::TimedOut => {}
WaitOutcome::Closed => return,
}
} else {
if send_event_batch(request_id, &batch, conn).is_err() {
return;
}
announced = false;
}
}
}
fn send_event_batch(
request_id: u64,
events: &[(Position, Event)],
conn: &ConnCtx,
) -> Result<(), ()> {
let mut batch = pb::ReadEvents::new();
let mut batch_bytes = 0usize;
for (position, event) in events {
batch_bytes += event.as_bytes().len();
batch
.events_mut()
.push(convert::sequenced_to_proto(*position, event.as_ref()));
if batch.events().len() >= conn.config.read_batch_events
|| batch_bytes >= conn.config.read_batch_bytes
{
let full = std::mem::replace(&mut batch, pb::ReadEvents::new());
conn.out_tx
.send(make_response(request_id, ResponseKind::ReadEvents(full)))
.map_err(|_| ())?;
batch_bytes = 0;
}
}
if !batch.events().is_empty() {
conn.out_tx
.send(make_response(request_id, ResponseKind::ReadEvents(batch)))
.map_err(|_| ())?;
}
Ok(())
}
fn writer_loop(
write_half: TcpStream,
out_rx: Receiver<pb::Response>,
shutdown: Option<TcpStream>,
alive: Arc<AtomicBool>,
max_frame_len: u32,
peer: Option<SocketAddr>,
) {
let mut writer = BufWriter::new(write_half);
let outcome = 'outer: loop {
let Ok(response) = out_rx.recv() else {
break Ok(());
};
if let Err(err) = write_frame(&mut writer, &response, max_frame_len) {
break Err(err);
}
while let Ok(response) = out_rx.try_recv() {
if let Err(err) = write_frame(&mut writer, &response, max_frame_len) {
break 'outer Err(err);
}
}
if let Err(err) = writer.flush() {
break Err(FrameError::Io(err));
}
};
if let Err(err) = outcome {
if alive.load(Ordering::Acquire) {
tracing::warn!(?peer, %err, "closing connection: writer failed");
} else {
tracing::debug!(?peer, %err, "writer ended after the connection was closed");
}
}
alive.store(false, Ordering::Release);
if let Some(stream) = shutdown {
let _ = stream.shutdown(Shutdown::Both);
}
}
fn pump_loop(
reply_rx: Receiver<AppendReply>,
out_tx: Sender<pb::Response>,
inflight: Arc<Semaphore>,
) {
while let Ok((request_id, result)) = reply_rx.recv() {
let response = match result {
Ok(range) => {
let mut ok = pb::AppendResponse::new();
ok.set_first(range.first.get());
ok.set_last(range.last.get());
make_response(request_id, ResponseKind::Append(ok))
}
Err(err) => make_response(
request_id,
ResponseKind::Error(convert::append_error_to_proto(&err)),
),
};
let sent = out_tx.send(response);
inflight.release();
if sent.is_err() {
break;
}
}
}
enum ResponseKind {
Append(pb::AppendResponse),
ReadEvents(pb::ReadEvents),
ReadEnd(pb::ReadEnd),
CaughtUp(pb::SubscribeCaughtUp),
Error(pb::ErrorResponse),
}
fn make_response(request_id: u64, kind: ResponseKind) -> pb::Response {
let mut response = pb::Response::new();
response.set_request_id(request_id);
match kind {
ResponseKind::Append(append) => response.set_append(append),
ResponseKind::ReadEvents(events) => response.set_read_events(events),
ResponseKind::ReadEnd(end) => response.set_read_end(end),
ResponseKind::CaughtUp(caught_up) => response.set_caught_up(caught_up),
ResponseKind::Error(error) => response.set_error(error),
}
response
}
pub(crate) struct ReadJob {
request_id: u64,
query: Query,
after: Position,
limit: Option<u64>,
conn: ConnCtx,
cancel: Arc<AtomicBool>,
cleanup: WorkerCleanup,
}
impl ReadJob {
fn run(self) {
let ReadJob {
request_id,
query,
after,
limit,
conn,
cancel,
cleanup,
} = self;
run_read(request_id, query, after, limit, &conn, &cancel);
drop(cleanup);
}
}
pub(crate) struct ReadPool {
tx: Sender<ReadJob>,
workers: Vec<thread::JoinHandle<()>>,
}
impl ReadPool {
pub(crate) fn new(size: usize) -> ReadPool {
let size = size.max(1);
let (tx, rx) = flume::unbounded::<ReadJob>();
let mut workers = Vec::with_capacity(size);
for _ in 0..size {
let rx = rx.clone();
let worker = thread::Builder::new()
.name("tephra-read-worker".to_string())
.spawn(move || {
while let Ok(job) = rx.recv() {
let _ = std::panic::catch_unwind(AssertUnwindSafe(move || job.run()));
}
})
.expect("spawn read worker thread");
workers.push(worker);
}
ReadPool { tx, workers }
}
pub(crate) fn sender(&self) -> Sender<ReadJob> {
self.tx.clone()
}
pub(crate) fn shutdown(self) {
drop(self.tx);
for worker in self.workers {
let _ = worker.join();
}
}
}
struct WorkerCleanup {
cancels: Arc<Mutex<HashMap<u64, Arc<AtomicBool>>>>,
sem: Option<Arc<Semaphore>>,
workers: WaitGroup,
request_id: u64,
}
impl Drop for WorkerCleanup {
fn drop(&mut self) {
self.cancels.lock().unwrap().remove(&self.request_id);
if let Some(sem) = &self.sem {
sem.release();
}
self.workers.done();
}
}
struct Semaphore {
permits: Mutex<usize>,
available: Condvar,
}
impl Semaphore {
fn new(permits: usize) -> Semaphore {
Semaphore {
permits: Mutex::new(permits.max(1)),
available: Condvar::new(),
}
}
fn acquire(&self) {
let mut permits = self.permits.lock().unwrap();
while *permits == 0 {
permits = self.available.wait(permits).unwrap();
}
*permits -= 1;
}
fn try_acquire(&self) -> bool {
let mut permits = self.permits.lock().unwrap();
if *permits == 0 {
false
} else {
*permits -= 1;
true
}
}
fn release(&self) {
*self.permits.lock().unwrap() += 1;
self.available.notify_one();
}
}
#[derive(Clone, Default)]
struct WaitGroup {
inner: Arc<(Mutex<usize>, Condvar)>,
}
impl WaitGroup {
fn add(&self) {
*self.inner.0.lock().unwrap() += 1;
}
fn done(&self) {
let mut count = self.inner.0.lock().unwrap();
*count -= 1;
if *count == 0 {
self.inner.1.notify_all();
}
}
fn wait(&self) {
let mut count = self.inner.0.lock().unwrap();
while *count > 0 {
count = self.inner.1.wait(count).unwrap();
}
}
}