use std::collections::HashMap;
use std::pin::Pin;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use futures_core::Stream;
use tokio::io::{AsyncWriteExt, BufWriter};
use tokio::net::tcp::{OwnedReadHalf, OwnedWriteHalf};
use tokio::net::{TcpStream, ToSocketAddrs};
use tokio::sync::{OwnedSemaphorePermit, Semaphore, mpsc, oneshot};
use tokio_stream::StreamExt;
use tephra_proto::convert as wire;
use tephra_proto::tephra as pb;
use tephra_proto::{DEFAULT_MAX_FRAME_LEN, read_frame_async, write_frame_async};
use super::{
AppendCondition, AppendResult, ClientError, Event, Position, Query, SequencedEvent, SubEvent,
UNATTRIBUTED_REQUEST_ID, event_to_pb, sequenced_from_pb, server_error,
};
#[derive(Clone, Copy, Debug)]
pub struct AsyncClientConfig {
pub max_frame_len: u32,
pub request_queue_depth: usize,
pub max_inflight_requests: usize,
}
impl Default for AsyncClientConfig {
fn default() -> Self {
AsyncClientConfig {
max_frame_len: DEFAULT_MAX_FRAME_LEN,
request_queue_depth: 256,
max_inflight_requests: 1024,
}
}
}
struct Shared {
next_id: AtomicU64,
requests: Mutex<HashMap<u64, Registered>>,
inflight: Arc<Semaphore>,
max_frame_len: u32,
}
impl Shared {
fn next_id(&self) -> u64 {
self.next_id.fetch_add(1, Ordering::Relaxed)
}
}
struct Registered {
sink: Sink,
_permit: OwnedSemaphorePermit,
}
enum Sink {
Append(oneshot::Sender<Result<AppendResult, ClientError>>),
Read(mpsc::UnboundedSender<ReadItem>),
Subscribe(mpsc::UnboundedSender<Result<SubEvent, ClientError>>),
}
enum ReadItem {
Event(SequencedEvent),
End(Position),
Err(ClientError),
}
#[derive(Clone)]
pub struct AsyncClient {
shared: Arc<Shared>,
out_tx: mpsc::Sender<pb::Request>,
}
impl AsyncClient {
pub async fn connect(addr: impl ToSocketAddrs) -> std::io::Result<AsyncClient> {
AsyncClient::connect_with(addr, AsyncClientConfig::default()).await
}
pub async fn connect_with(
addr: impl ToSocketAddrs,
config: AsyncClientConfig,
) -> std::io::Result<AsyncClient> {
let stream = TcpStream::connect(addr).await?;
stream.set_nodelay(true)?;
let (read_half, write_half) = stream.into_split();
let (out_tx, out_rx) = mpsc::channel(config.request_queue_depth.max(1));
let shared = Arc::new(Shared {
next_id: AtomicU64::new(1),
requests: Mutex::new(HashMap::new()),
inflight: Arc::new(Semaphore::new(config.max_inflight_requests.max(1))),
max_frame_len: config.max_frame_len,
});
tokio::spawn(reader_task(read_half, Arc::clone(&shared)));
tokio::spawn(writer_task(write_half, out_rx, shared.max_frame_len));
Ok(AsyncClient { shared, out_tx })
}
async fn acquire_inflight(&self) -> OwnedSemaphorePermit {
Arc::clone(&self.shared.inflight)
.acquire_owned()
.await
.expect("inflight semaphore is never closed")
}
pub async fn append(
&self,
events: impl IntoIterator<Item = Event>,
condition: Option<AppendCondition>,
) -> Result<AppendResult, ClientError> {
let id = self.shared.next_id();
let mut append = pb::AppendRequest::new();
for event in events {
append.events_mut().push(event_to_pb(&event));
}
if let Some(condition) = condition {
append.set_condition(wire::condition_to_pb(&condition));
}
let mut request = pb::Request::new();
request.set_request_id(id);
request.set_append(append);
let _permit = self.acquire_inflight().await;
let (tx, rx) = oneshot::channel();
self.shared.requests.lock().unwrap().insert(
id,
Registered {
sink: Sink::Append(tx),
_permit,
},
);
if self.out_tx.send(request).await.is_err() {
self.shared.requests.lock().unwrap().remove(&id);
return Err(ClientError::UnexpectedEof);
}
rx.await.unwrap_or(Err(ClientError::UnexpectedEof))
}
pub async fn read(&self, query: Query, after: Position, limit: Option<u64>) -> ReadStream {
let id = self.shared.next_id();
let mut read = pb::ReadRequest::new();
read.set_query(wire::query_to_pb(&query));
read.set_after(after.get());
if let Some(limit) = limit {
read.set_limit(limit);
}
let mut request = pb::Request::new();
request.set_request_id(id);
request.set_read(read);
let _permit = self.acquire_inflight().await;
let (tx, rx) = mpsc::unbounded_channel();
self.shared.requests.lock().unwrap().insert(
id,
Registered {
sink: Sink::Read(tx),
_permit,
},
);
if self.out_tx.send(request).await.is_err() {
if let Some(Registered {
sink: Sink::Read(tx),
..
}) = self.shared.requests.lock().unwrap().remove(&id)
{
let _ = tx.send(ReadItem::Err(ClientError::UnexpectedEof));
}
}
ReadStream {
shared: Arc::clone(&self.shared),
out_tx: self.out_tx.clone(),
id,
rx,
watermark: None,
done: false,
}
}
pub async fn read_all(
&self,
query: Query,
after: Position,
limit: Option<u64>,
) -> Result<(Vec<SequencedEvent>, Position), ClientError> {
let mut stream = self.read(query, after, limit).await;
let mut events = Vec::new();
while let Some(item) = stream.next().await {
events.push(item?);
}
let watermark = stream
.watermark()
.ok_or_else(|| ClientError::Protocol("read ended without a watermark".to_string()))?;
Ok((events, watermark))
}
pub async fn subscribe(&self, query: Query, after: Position) -> SubscribeStream {
let id = self.shared.next_id();
let mut subscribe = pb::SubscribeRequest::new();
subscribe.set_query(wire::query_to_pb(&query));
subscribe.set_after(after.get());
let mut request = pb::Request::new();
request.set_request_id(id);
request.set_subscribe(subscribe);
let _permit = self.acquire_inflight().await;
let (tx, rx) = mpsc::unbounded_channel();
self.shared.requests.lock().unwrap().insert(
id,
Registered {
sink: Sink::Subscribe(tx),
_permit,
},
);
if self.out_tx.send(request).await.is_err() {
if let Some(Registered {
sink: Sink::Subscribe(tx),
..
}) = self.shared.requests.lock().unwrap().remove(&id)
{
let _ = tx.send(Err(ClientError::UnexpectedEof));
}
}
SubscribeStream {
shared: Arc::clone(&self.shared),
out_tx: self.out_tx.clone(),
id,
rx,
done: false,
}
}
}
async fn writer_task(
write_half: OwnedWriteHalf,
mut out_rx: mpsc::Receiver<pb::Request>,
max_frame_len: u32,
) {
let mut writer = BufWriter::new(write_half);
while let Some(request) = out_rx.recv().await {
if write_frame_async(&mut writer, &request, max_frame_len)
.await
.is_err()
{
break;
}
while let Ok(request) = out_rx.try_recv() {
if write_frame_async(&mut writer, &request, max_frame_len)
.await
.is_err()
{
return;
}
}
if writer.flush().await.is_err() {
break;
}
}
}
async fn reader_task(mut read_half: OwnedReadHalf, shared: Arc<Shared>) {
let max = shared.max_frame_len;
let mut last_error: Option<String> = None;
loop {
match read_frame_async::<pb::Response, _>(&mut read_half, max).await {
Ok(Some(response)) => {
if response.request_id() == UNATTRIBUTED_REQUEST_ID {
if let pb::response::KindOneof::Error(error) = response.kind() {
last_error = Some(
error
.message()
.to_str()
.unwrap_or("server error")
.to_string(),
);
}
continue;
}
route(response, &shared);
}
Ok(None) => break,
Err(err) => {
last_error = Some(format!("connection error: {err}"));
break;
}
}
}
let reason = last_error.unwrap_or_else(|| "server closed the connection".to_string());
fail_all(&shared, &reason);
}
fn route(response: pb::Response, shared: &Shared) {
let id = response.request_id();
let mut map = shared.requests.lock().unwrap();
let Some(Registered { sink, _permit }) = map.remove(&id) else {
return;
};
match sink {
Sink::Append(tx) => {
let result = match response.kind() {
pb::response::KindOneof::Append(append) => Ok(AppendResult {
first: Position::new(append.first()),
last: Position::new(append.last()),
}),
pb::response::KindOneof::Error(error) => Err(server_error(error)),
other => Err(ClientError::Protocol(format!(
"unexpected response to append: {other:?}"
))),
};
let _ = tx.send(result);
}
Sink::Read(tx) => {
if deliver_read(&tx, response) {
map.insert(
id,
Registered {
sink: Sink::Read(tx),
_permit,
},
);
}
}
Sink::Subscribe(tx) => {
if deliver_subscribe(&tx, response) {
map.insert(
id,
Registered {
sink: Sink::Subscribe(tx),
_permit,
},
);
}
}
}
}
fn deliver_read(tx: &mpsc::UnboundedSender<ReadItem>, response: pb::Response) -> bool {
match response.kind() {
pb::response::KindOneof::ReadEvents(events) => {
for view in events.events().iter() {
match sequenced_from_pb(view) {
Ok(event) => {
if tx.send(ReadItem::Event(event)).is_err() {
return false;
}
}
Err(err) => {
let _ = tx.send(ReadItem::Err(err));
return false;
}
}
}
true
}
pb::response::KindOneof::ReadEnd(end) => {
let _ = tx.send(ReadItem::End(Position::new(end.watermark())));
false
}
pb::response::KindOneof::Error(error) => {
let _ = tx.send(ReadItem::Err(server_error(error)));
false
}
other => {
let _ = tx.send(ReadItem::Err(ClientError::Protocol(format!(
"unexpected response during read: {other:?}"
))));
false
}
}
}
fn deliver_subscribe(
tx: &mpsc::UnboundedSender<Result<SubEvent, ClientError>>,
response: pb::Response,
) -> bool {
match response.kind() {
pb::response::KindOneof::ReadEvents(events) => {
for view in events.events().iter() {
match sequenced_from_pb(view) {
Ok(event) => {
if tx.send(Ok(SubEvent::Event(event))).is_err() {
return false;
}
}
Err(err) => {
let _ = tx.send(Err(err));
return false;
}
}
}
true
}
pb::response::KindOneof::CaughtUp(caught_up) => tx
.send(Ok(SubEvent::CaughtUp(Position::new(caught_up.watermark()))))
.is_ok(),
pb::response::KindOneof::Error(error) => {
let _ = tx.send(Err(server_error(error)));
false
}
other => {
let _ = tx.send(Err(ClientError::Protocol(format!(
"unexpected response during subscribe: {other:?}"
))));
false
}
}
}
fn fail_all(shared: &Shared, reason: &str) {
let mut map = shared.requests.lock().unwrap();
for (_id, Registered { sink, _permit }) in map.drain() {
match sink {
Sink::Append(tx) => {
let _ = tx.send(Err(ClientError::Protocol(reason.to_string())));
}
Sink::Read(tx) => {
let _ = tx.send(ReadItem::Err(ClientError::Protocol(reason.to_string())));
}
Sink::Subscribe(tx) => {
let _ = tx.send(Err(ClientError::Protocol(reason.to_string())));
}
}
}
}
fn send_cancel(out_tx: &mpsc::Sender<pb::Request>, shared: &Shared, target: u64) {
let mut cancel = pb::CancelRequest::new();
cancel.set_target(target);
let mut request = pb::Request::new();
request.set_request_id(shared.next_id());
request.set_cancel(cancel);
let _ = out_tx.try_send(request);
}
pub struct ReadStream {
shared: Arc<Shared>,
out_tx: mpsc::Sender<pb::Request>,
id: u64,
rx: mpsc::UnboundedReceiver<ReadItem>,
watermark: Option<Position>,
done: bool,
}
impl ReadStream {
pub fn watermark(&self) -> Option<Position> {
self.watermark
}
}
impl Stream for ReadStream {
type Item = Result<SequencedEvent, ClientError>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
if this.done {
return Poll::Ready(None);
}
match this.rx.poll_recv(cx) {
Poll::Ready(Some(ReadItem::Event(event))) => Poll::Ready(Some(Ok(event))),
Poll::Ready(Some(ReadItem::End(watermark))) => {
this.watermark = Some(watermark);
this.done = true;
Poll::Ready(None)
}
Poll::Ready(Some(ReadItem::Err(err))) => {
this.done = true;
Poll::Ready(Some(Err(err)))
}
Poll::Ready(None) => {
this.done = true;
Poll::Ready(None)
}
Poll::Pending => Poll::Pending,
}
}
}
impl Drop for ReadStream {
fn drop(&mut self) {
if !self.done {
send_cancel(&self.out_tx, &self.shared, self.id);
}
self.shared.requests.lock().unwrap().remove(&self.id);
}
}
pub struct SubscribeStream {
shared: Arc<Shared>,
out_tx: mpsc::Sender<pb::Request>,
id: u64,
rx: mpsc::UnboundedReceiver<Result<SubEvent, ClientError>>,
done: bool,
}
impl Stream for SubscribeStream {
type Item = Result<SubEvent, ClientError>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
if this.done {
return Poll::Ready(None);
}
match this.rx.poll_recv(cx) {
Poll::Ready(Some(Ok(event))) => Poll::Ready(Some(Ok(event))),
Poll::Ready(Some(Err(err))) => {
this.done = true;
Poll::Ready(Some(Err(err)))
}
Poll::Ready(None) => {
this.done = true;
Poll::Ready(None)
}
Poll::Pending => Poll::Pending,
}
}
}
impl Drop for SubscribeStream {
fn drop(&mut self) {
if !self.done {
send_cancel(&self.out_tx, &self.shared, self.id);
}
self.shared.requests.lock().unwrap().remove(&self.id);
}
}