use crate::api::error::ApiError;
use crate::stream::StreamEvent;
use futures::Stream;
use std::pin::Pin;
use std::task::{Context, Poll};
use std::time::{Duration, Instant};
#[derive(Debug, Clone)]
pub struct HeartbeatData {
pub elapsed: Duration,
pub is_timeout: bool,
}
pub type HeartbeatCallback = Box<dyn Fn(HeartbeatData) + Send + Sync>;
pub struct HeartbeatConfig {
heartbeat_interval: Duration,
timeout: Duration,
on_heartbeat: HeartbeatCallback,
}
impl std::fmt::Debug for HeartbeatConfig {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("HeartbeatConfig")
.field("heartbeat_interval", &self.heartbeat_interval)
.field("timeout", &self.timeout)
.finish_non_exhaustive()
}
}
impl HeartbeatConfig {
#[must_use]
pub fn new(
heartbeat_interval: Duration,
timeout: Duration,
on_heartbeat: HeartbeatCallback,
) -> Self {
Self {
heartbeat_interval,
timeout,
on_heartbeat,
}
}
#[must_use]
pub fn heartbeat_interval(&self) -> Duration {
self.heartbeat_interval
}
#[must_use]
pub fn timeout(&self) -> Duration {
self.timeout
}
}
pub struct HeartbeatStream<S> {
inner: S,
config: HeartbeatConfig,
last_heartbeat: Instant,
start: Instant,
timeout_sleep: std::pin::Pin<Box<tokio::time::Sleep>>,
}
impl<S> HeartbeatStream<S> {
pub fn new(inner: S, config: HeartbeatConfig) -> Self {
const THIRTY_YEARS_SECS: u64 = 86400 * 365 * 30;
let now = Instant::now();
let far_future = || {
Instant::now()
.checked_add(Duration::from_secs(THIRTY_YEARS_SECS))
.unwrap_or(Instant::now())
};
let deadline = now.checked_add(config.timeout).unwrap_or_else(far_future);
let timeout_sleep = Box::pin(tokio::time::sleep_until(tokio::time::Instant::from_std(
deadline,
)));
Self {
inner,
config,
last_heartbeat: now,
start: now,
timeout_sleep,
}
}
}
impl<S> Stream for HeartbeatStream<S>
where
S: Stream<Item = Result<StreamEvent, ApiError>> + Unpin,
{
type Item = Result<StreamEvent, ApiError>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.get_mut();
if this.last_heartbeat.elapsed() >= this.config.heartbeat_interval {
let elapsed = this.start.elapsed();
let data = HeartbeatData {
elapsed,
is_timeout: elapsed > this.config.timeout,
};
(this.config.on_heartbeat)(data);
this.last_heartbeat = Instant::now();
}
if this.timeout_sleep.as_mut().poll(cx).is_ready()
|| this.start.elapsed() > this.config.timeout
{
return Poll::Ready(Some(Err(ApiError::Api(format!(
"Stream timeout after {}s",
this.config.timeout.as_secs()
)))));
}
Pin::new(&mut this.inner).poll_next(cx)
}
}
#[cfg(test)]
mod tests {
use super::*;
use futures::StreamExt;
struct VecStream {
items: Vec<Result<StreamEvent, ApiError>>,
}
impl Stream for VecStream {
type Item = Result<StreamEvent, ApiError>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Ready(self.get_mut().items.pop())
}
}
fn make_config(
callbacks: &std::sync::Arc<std::sync::Mutex<Vec<HeartbeatData>>>,
) -> HeartbeatConfig {
let cb = callbacks.clone();
HeartbeatConfig::new(
Duration::from_millis(10),
Duration::from_secs(60),
Box::new(move |data: HeartbeatData| {
cb.lock().unwrap().push(data);
}),
)
}
#[test]
fn heartbeat_data_fields() {
let data = HeartbeatData {
elapsed: Duration::from_secs(30),
is_timeout: true,
};
assert_eq!(data.elapsed, Duration::from_secs(30));
assert!(data.is_timeout);
}
#[test]
fn config_accessors() {
let config = HeartbeatConfig::new(
Duration::from_secs(15),
Duration::from_secs(300),
Box::new(|_| {}),
);
assert_eq!(config.heartbeat_interval(), Duration::from_secs(15));
assert_eq!(config.timeout(), Duration::from_secs(300));
}
#[test]
fn config_debug() {
let config = HeartbeatConfig::new(
Duration::from_secs(30),
Duration::from_secs(600),
Box::new(|_| {}),
);
let debug = format!("{config:?}");
assert!(debug.contains("HeartbeatConfig"));
assert!(debug.contains("heartbeat_interval"));
assert!(debug.contains("timeout"));
}
#[tokio::test]
async fn passes_through_events() {
let callbacks: std::sync::Arc<std::sync::Mutex<Vec<HeartbeatData>>> =
std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let config = make_config(&callbacks);
let inner = VecStream {
items: vec![Ok(StreamEvent::Ping), Ok(StreamEvent::Ping)],
};
let mut stream = HeartbeatStream::new(inner, config);
let first = stream.next().await;
assert!(first.is_some());
let second = stream.next().await;
assert!(second.is_some());
let third = stream.next().await;
assert!(third.is_none());
}
#[tokio::test]
async fn passes_through_errors() {
let callbacks: std::sync::Arc<std::sync::Mutex<Vec<HeartbeatData>>> =
std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let config = make_config(&callbacks);
let inner = VecStream {
items: vec![Err(ApiError::Api("test error".to_string()))],
};
let mut stream = HeartbeatStream::new(inner, config);
let result = stream.next().await;
assert!(matches!(result, Some(Err(ApiError::Api(_)))));
}
#[tokio::test]
async fn fires_heartbeat_on_interval_sync() {
let callbacks: std::sync::Arc<std::sync::Mutex<Vec<HeartbeatData>>> =
std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let config = make_config(&callbacks);
let inner = VecStream {
items: vec![Ok(StreamEvent::Ping)],
};
let mut stream = HeartbeatStream::new(inner, config);
stream.last_heartbeat = Instant::now().checked_sub(Duration::from_secs(1)).unwrap();
let waker = futures::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let result = Pin::new(&mut stream).poll_next(&mut cx);
assert!(matches!(result, Poll::Ready(Some(Ok(StreamEvent::Ping)))));
let cbs = callbacks.lock().unwrap();
assert_eq!(cbs.len(), 1);
assert!(cbs[0].elapsed > Duration::ZERO);
}
#[tokio::test]
async fn timeout_returns_error_sync() {
let callbacks: std::sync::Arc<std::sync::Mutex<Vec<HeartbeatData>>> =
std::sync::Arc::new(std::sync::Mutex::new(Vec::new()));
let config = HeartbeatConfig::new(
Duration::from_millis(10),
Duration::from_millis(1),
Box::new(move |data: HeartbeatData| {
callbacks.lock().unwrap().push(data);
}),
);
let inner = VecStream { items: vec![] };
let mut stream = HeartbeatStream::new(inner, config);
stream.start = Instant::now().checked_sub(Duration::from_secs(10)).unwrap();
let waker = futures::task::noop_waker();
let mut cx = Context::from_waker(&waker);
let result = Pin::new(&mut stream).poll_next(&mut cx);
assert!(
matches!(result, Poll::Ready(Some(Err(ApiError::Api(msg)))) if msg.contains("timeout"))
);
}
}