use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use futures::Stream;
use ruststream::Subscriber;
use sea_streamer_file::{FileConsumer, FileErr};
use sea_streamer_types::{Consumer as _, SeqPos, StreamErr, Timestamp};
use tokio::sync::{mpsc, oneshot};
use crate::error::{SeaFileError, box_err};
use crate::message::{FilePosition, SeaMessage};
const CHANNEL_CAPACITY: usize = 64;
pub(crate) struct SeekCmd {
position: FilePosition,
done: oneshot::Sender<Result<(), SeaFileError>>,
}
pub(crate) struct Stamped {
epoch: u64,
item: Option<Result<SeaMessage, SeaFileError>>,
}
pub struct FileSubscriber {
stream: String,
rx: mpsc::Receiver<Stamped>,
cmd: mpsc::UnboundedSender<SeekCmd>,
epoch: Arc<AtomicU64>,
}
impl std::fmt::Debug for FileSubscriber {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FileSubscriber")
.field("stream", &self.stream)
.finish_non_exhaustive()
}
}
impl FileSubscriber {
#[must_use]
pub fn stream_key(&self) -> &str {
&self.stream
}
pub(crate) fn spawn(stream: String, consumer: FileConsumer, replay: bool) -> Self {
let (out_tx, out_rx) = mpsc::channel(CHANNEL_CAPACITY);
let (cmd_tx, cmd_rx) = mpsc::unbounded_channel();
let epoch = Arc::new(AtomicU64::new(0));
tokio::spawn(drive(
consumer,
out_tx,
cmd_rx,
stream.clone(),
replay,
Arc::clone(&epoch),
));
Self {
stream,
rx: out_rx,
cmd: cmd_tx,
epoch,
}
}
}
impl Subscriber for FileSubscriber {
type Message = SeaMessage;
type Error = SeaFileError;
fn stream(&mut self) -> impl Stream<Item = Result<SeaMessage, SeaFileError>> + Send + '_ {
futures::stream::poll_fn(move |cx| {
loop {
match self.rx.poll_recv(cx) {
std::task::Poll::Ready(Some(stamped)) => {
if stamped.epoch == self.epoch.load(Ordering::Acquire) {
return std::task::Poll::Ready(stamped.item);
}
}
std::task::Poll::Ready(None) => return std::task::Poll::Ready(None),
std::task::Poll::Pending => return std::task::Poll::Pending,
}
}
})
}
}
#[derive(Clone)]
pub struct FileSeeker {
cmd: mpsc::UnboundedSender<SeekCmd>,
epoch: Arc<AtomicU64>,
stream: String,
}
impl std::fmt::Debug for FileSeeker {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("FileSeeker")
.field("stream", &self.stream)
.finish_non_exhaustive()
}
}
impl ruststream::Seeker for FileSeeker {
type Position = FilePosition;
type Error = SeaFileError;
async fn seek(&self, to: FilePosition) -> Result<(), SeaFileError> {
self.epoch.fetch_add(1, Ordering::Release);
let (done, wait) = oneshot::channel();
self.cmd
.send(SeekCmd { position: to, done })
.map_err(|_| SeaFileError::Seek {
stream: self.stream.clone(),
source: Box::from("the subscription's driver task has shut down"),
})?;
wait.await.map_err(|_| SeaFileError::Seek {
stream: self.stream.clone(),
source: Box::from("the subscription's driver task has shut down"),
})?
}
}
impl ruststream::Seekable for FileSubscriber {
type Seeker = FileSeeker;
fn seeker(&self) -> FileSeeker {
FileSeeker {
cmd: self.cmd.clone(),
epoch: Arc::clone(&self.epoch),
stream: self.stream.clone(),
}
}
}
fn is_clean_end(err: &StreamErr<FileErr>) -> bool {
matches!(
err,
StreamErr::Backend(FileErr::StreamEnded | FileErr::NotEnoughBytes)
)
}
async fn drive(
mut consumer: FileConsumer,
out: mpsc::Sender<Stamped>,
mut cmd_rx: mpsc::UnboundedReceiver<SeekCmd>,
stream: String,
replay: bool,
epoch: Arc<AtomicU64>,
) {
loop {
let current = epoch.load(Ordering::Acquire);
tokio::select! {
biased;
cmd = cmd_rx.recv() => {
let Some(SeekCmd { position, done }) = cmd else { break };
let result = match position {
FilePosition::Beginning => consumer.rewind(SeqPos::Beginning).await,
FilePosition::End => consumer.rewind(SeqPos::End).await,
FilePosition::Sequence(sequence) => {
consumer.rewind(SeqPos::At(sequence)).await
}
FilePosition::Timestamp(millis) => {
let nanos = i128::from(millis) * 1_000_000;
match Timestamp::from_unix_timestamp_nanos(nanos) {
Ok(timestamp) => consumer.seek(timestamp).await,
Err(err) => {
let _ = done.send(Err(SeaFileError::Invalid(format!(
"'{millis}' is not a valid timestamp: {err}"
))));
continue;
}
}
}
};
let _ = done.send(result.map_err(|e| SeaFileError::Seek {
stream: stream.clone(),
source: box_err(e),
}));
}
() = out.closed() => break,
next = consumer.next() => {
match next {
Ok(message) => {
let item = Stamped {
epoch: current,
item: Some(Ok(SeaMessage::new(&message))),
};
if out.send(item).await.is_err() {
break;
}
}
Err(err) if is_clean_end(&err) => {
if replay {
let _ = out.send(Stamped { epoch: current, item: None }).await;
} else {
let _ = out
.send(Stamped {
epoch: current,
item: Some(Err(SeaFileError::Receive {
stream: stream.clone(),
source: box_err(err),
})),
})
.await;
}
break;
}
Err(err) => {
let _ = out
.send(Stamped {
epoch: current,
item: Some(Err(SeaFileError::Receive {
stream: stream.clone(),
source: box_err(err),
})),
})
.await;
break;
}
}
}
}
}
}