use std::marker::PhantomData;
use std::pin::Pin;
use std::sync::Arc;
use std::task::Waker;
use std::task::{Context, Poll};
use std::time::{Duration, Instant};
use futures::ready;
use futures::{Sink, Stream};
use log::debug;
use crate::errors::StreamError;
use crate::packet::{Operation, Packet, Protocol};
use super::waker::WakerProxy;
pub struct HeartbeatStream<T, E> {
stream: T,
tx_waker: Arc<WakerProxy>,
last_hb: Option<Instant>,
__marker: PhantomData<E>,
}
impl<T: Unpin, E> Unpin for HeartbeatStream<T, E> {}
impl<T, E> HeartbeatStream<T, E> {
pub fn new(stream: T) -> Self {
Self {
stream,
tx_waker: Arc::new(Default::default()),
last_hb: None,
__marker: PhantomData,
}
}
fn with_context<F, U>(&mut self, f: F) -> U
where
F: FnOnce(&mut Context<'_>, &mut T) -> U,
{
let waker = Waker::from(self.tx_waker.clone());
let mut cx = Context::from_waker(&waker);
f(&mut cx, &mut self.stream)
}
}
impl<T, E> Stream for HeartbeatStream<T, E>
where
T: Stream<Item = Result<Packet, StreamError<E>>> + Sink<Packet, Error = StreamError<E>> + Unpin,
E: std::error::Error,
{
type Item = Result<Packet, StreamError<E>>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.tx_waker.rx(cx.waker());
ready!(self.with_context(|cx, s| Pin::new(s).poll_ready(cx)))?;
let now = Instant::now();
let need_hb = self
.last_hb
.map_or(true, |last_hb| now - last_hb >= Duration::from_secs(30));
if need_hb {
debug!("sending heartbeat");
self.as_mut()
.start_send(Packet::new(Operation::HeartBeat, Protocol::Json, vec![]))?;
self.last_hb = Some(now);
#[cfg(feature = "tokio")]
{
let waker = cx.waker().clone();
tokio1::spawn(async {
tokio1::time::sleep(Duration::from_secs(30)).await;
waker.wake();
});
}
#[cfg(feature = "async-std")]
{
let waker = cx.waker().clone();
async_std1::task::spawn(async {
async_std1::task::sleep(Duration::from_secs(30)).await;
waker.wake();
});
}
ready!(self.with_context(|cx, s| Pin::new(s).poll_flush(cx)))?;
}
Pin::new(&mut self.stream).poll_next(cx)
}
}
impl<T, E> Sink<Packet> for HeartbeatStream<T, E>
where
T: Sink<Packet, Error = StreamError<E>> + Unpin,
E: std::error::Error,
{
type Error = StreamError<E>;
fn poll_ready(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.tx_waker.tx(cx.waker());
self.with_context(|cx, s| Pin::new(s).poll_ready(cx))
}
fn start_send(mut self: Pin<&mut Self>, item: Packet) -> Result<(), Self::Error> {
Pin::new(&mut self.stream).start_send(item)
}
fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.tx_waker.tx(cx.waker());
self.with_context(|cx, s| Pin::new(s).poll_flush(cx))
}
fn poll_close(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Result<(), Self::Error>> {
self.tx_waker.tx(cx.waker());
self.with_context(|cx, s| Pin::new(s).poll_close(cx))
}
}