use std::io::{BufReader, BufWriter, Write};
use std::net::TcpStream;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tephra::read::WaitOutcome;
use tephra::writer::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;
pub(crate) fn serve_connection(
stream: TcpStream,
handle: WriteHandle,
config: ServerConfig,
running: Arc<AtomicBool>,
) {
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_stream = match stream.try_clone() {
Ok(clone) => clone,
Err(err) => {
tracing::warn!(?peer, %err, "failed to clone connection stream");
return;
}
};
let mut reader = BufReader::new(read_stream);
let mut writer = BufWriter::new(stream);
loop {
let request = match read_frame::<pb::Request, _>(&mut reader, config.max_frame_len) {
Ok(Some(request)) => request,
Ok(None) => 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 _ = send(
&mut writer,
0,
ResponseKind::Error(error),
config.max_frame_len,
);
}
tracing::debug!(?peer, %err, "connection read ended");
break;
}
};
if let Err(err) = dispatch(&request, &handle, &config, &running, &mut writer) {
tracing::debug!(?peer, %err, "connection write ended");
break;
}
}
}
fn dispatch(
request: &pb::Request,
handle: &WriteHandle,
config: &ServerConfig,
running: &AtomicBool,
writer: &mut BufWriter<TcpStream>,
) -> Result<(), FrameError> {
let request_id = request.request_id();
match request.kind() {
pb::request::KindOneof::Append(append) => {
handle_append(request_id, append, handle, config, writer)
}
pb::request::KindOneof::Read(read) => handle_read(request_id, read, handle, config, writer),
pb::request::KindOneof::Subscribe(subscribe) => {
handle_subscribe(request_id, subscribe, handle, config, running, writer)
}
_ => {
let error = convert::bad_request("request has no append, read, or subscribe set");
send(
writer,
request_id,
ResponseKind::Error(error),
config.max_frame_len,
)
}
}
}
fn handle_append(
request_id: u64,
append: pb::AppendRequestView<'_>,
handle: &WriteHandle,
config: &ServerConfig,
writer: &mut BufWriter<TcpStream>,
) -> Result<(), FrameError> {
let events = match convert::events_from_proto(append) {
Ok(events) => events,
Err(err) => {
let error = convert::bad_request(err);
return send(
writer,
request_id,
ResponseKind::Error(error),
config.max_frame_len,
);
}
};
let condition = match append.condition_opt() {
Some(condition) => match convert::condition_from_proto(condition) {
Ok(condition) => Some(condition),
Err(err) => {
let error = convert::bad_request(err);
return send(
writer,
request_id,
ResponseKind::Error(error),
config.max_frame_len,
);
}
},
None => None,
};
match handle.append(events, condition) {
Ok(range) => {
let mut ok = pb::AppendResponse::new();
ok.set_first(range.first.get());
ok.set_last(range.last.get());
send(
writer,
request_id,
ResponseKind::Append(ok),
config.max_frame_len,
)
}
Err(err) => {
let error = convert::append_error_to_proto(&err);
send(
writer,
request_id,
ResponseKind::Error(error),
config.max_frame_len,
)
}
}
}
fn handle_read(
request_id: u64,
read: pb::ReadRequestView<'_>,
handle: &WriteHandle,
config: &ServerConfig,
writer: &mut BufWriter<TcpStream>,
) -> Result<(), FrameError> {
let query = match convert::query_from_proto(read.query()) {
Ok(query) => query,
Err(err) => {
let error = convert::bad_request(err);
return send(
writer,
request_id,
ResponseKind::Error(error),
config.max_frame_len,
);
}
};
let after = Position::new(read.after());
let mut reads = handle.read(query, after);
let watermark = reads.watermark();
let mut batch = pb::ReadEvents::new();
let mut batch_bytes = 0usize;
while let Some(item) = reads.next() {
let sequenced = match item {
Ok(sequenced) => sequenced,
Err(err) => {
let error = convert::internal_error(err);
return send(
writer,
request_id,
ResponseKind::Error(error),
config.max_frame_len,
);
}
};
batch_bytes += sequenced.event.as_bytes().len();
batch.events_mut().push(convert::sequenced_to_proto(
sequenced.position,
sequenced.event,
));
if batch.events().len() >= config.read_batch_events
|| batch_bytes >= config.read_batch_bytes
{
send(
writer,
request_id,
ResponseKind::ReadEvents(batch),
config.max_frame_len,
)?;
batch = pb::ReadEvents::new();
batch_bytes = 0;
}
}
if !batch.events().is_empty() {
send(
writer,
request_id,
ResponseKind::ReadEvents(batch),
config.max_frame_len,
)?;
}
let mut end = pb::ReadEnd::new();
end.set_watermark(watermark.get());
send(
writer,
request_id,
ResponseKind::ReadEnd(end),
config.max_frame_len,
)
}
fn handle_subscribe(
request_id: u64,
subscribe: pb::SubscribeRequestView<'_>,
handle: &WriteHandle,
config: &ServerConfig,
running: &AtomicBool,
writer: &mut BufWriter<TcpStream>,
) -> Result<(), FrameError> {
let query = match convert::query_from_proto(subscribe.query()) {
Ok(query) => query,
Err(err) => {
let error = convert::bad_request(err);
return send(
writer,
request_id,
ResponseKind::Error(error),
config.max_frame_len,
);
}
};
let after = Position::new(subscribe.after());
let mut sub = handle.subscribe(query, after);
let mut announced = false;
loop {
if !running.load(Ordering::Acquire) {
return Ok(());
}
let batch = match sub.poll_batch() {
Ok(batch) => batch,
Err(err) => {
let error = convert::internal_error(err);
return send(
writer,
request_id,
ResponseKind::Error(error),
config.max_frame_len,
);
}
};
if batch.is_empty() {
if !announced {
let mut caught_up = pb::SubscribeCaughtUp::new();
caught_up.set_watermark(sub.position().get());
send(
writer,
request_id,
ResponseKind::CaughtUp(caught_up),
config.max_frame_len,
)?;
announced = true;
}
match sub.wait_timeout(config.subscribe_wait_tick) {
WaitOutcome::Advanced | WaitOutcome::TimedOut => {}
WaitOutcome::Closed => return Ok(()),
}
} else {
send_event_batch(request_id, &batch, config, writer)?;
announced = false;
}
}
}
fn send_event_batch(
request_id: u64,
events: &[(Position, Event)],
config: &ServerConfig,
writer: &mut BufWriter<TcpStream>,
) -> Result<(), FrameError> {
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() >= config.read_batch_events
|| batch_bytes >= config.read_batch_bytes
{
send(
writer,
request_id,
ResponseKind::ReadEvents(batch),
config.max_frame_len,
)?;
batch = pb::ReadEvents::new();
batch_bytes = 0;
}
}
if !batch.events().is_empty() {
send(
writer,
request_id,
ResponseKind::ReadEvents(batch),
config.max_frame_len,
)?;
}
Ok(())
}
enum ResponseKind {
Append(pb::AppendResponse),
ReadEvents(pb::ReadEvents),
ReadEnd(pb::ReadEnd),
CaughtUp(pb::SubscribeCaughtUp),
Error(pb::ErrorResponse),
}
fn send(
writer: &mut BufWriter<TcpStream>,
request_id: u64,
kind: ResponseKind,
max_frame_len: u32,
) -> Result<(), FrameError> {
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),
}
write_frame(writer, &response, max_frame_len)?;
writer.flush()?;
Ok(())
}