use std::future::Future;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex as StdMutex};
use std::time::Duration;
use tokio::sync::{mpsc, oneshot, watch};
use tokio_stream::wrappers::ReceiverStream;
use tokio_util::sync::CancellationToken;
use tonic::transport::Channel;
use crate::signet::v1::acquire_restart_lock_response::MessageType;
use crate::signet::v1::secrets_service_client::SecretsServiceClient;
use crate::signet::v1::watch_service_bundle_response::EventType;
use crate::signet::v1::{
AcquireRestartLockRequest, AcquireRestartLockResponse, WatchServiceBundleRequest,
WatchServiceBundleResponse,
};
const WATCH_BACKOFF_MIN: Duration = Duration::from_secs(1);
const WATCH_BACKOFF_MAX: Duration = Duration::from_secs(30);
#[derive(Debug, Clone, thiserror::Error)]
pub enum RestartError {
#[error("signet: ttl must be > 0, got {0:?}")]
InvalidTtl(Duration),
#[error("signet: open AcquireRestartLock stream: {0}")]
OpenStream(String),
#[error("signet: send AcquireRestartLockRequest: {0}")]
Send(String),
#[error("signet: receive AcquireRestartLockResponse: {0}")]
Recv(String),
#[error("signet: AcquireRestartLock stream closed before the lock was acquired")]
ClosedBeforeAcquired,
#[error("signet: restart lock lost: {0}")]
LockLost(String),
#[error("signet: open WatchServiceBundle stream: {0}")]
WatchOpen(String),
#[error("signet: watch_bundle stopped before a restart could be coordinated")]
WatchStopped,
#[error("signet: cancelled")]
Cancelled,
}
trait LockStream: Send {
fn send(
&mut self,
req: AcquireRestartLockRequest,
) -> impl Future<Output = Result<(), RestartError>> + Send;
fn recv(
&mut self,
) -> impl Future<Output = Result<Option<AcquireRestartLockResponse>, RestartError>> + Send;
fn close_send(&mut self) -> impl Future<Output = Result<(), RestartError>> + Send;
}
struct RealLockStream {
tx: Option<mpsc::Sender<AcquireRestartLockRequest>>,
inbound: tonic::codec::Streaming<AcquireRestartLockResponse>,
}
impl LockStream for RealLockStream {
async fn send(&mut self, req: AcquireRestartLockRequest) -> Result<(), RestartError> {
match &self.tx {
Some(tx) => tx
.send(req)
.await
.map_err(|_| RestartError::Send("stream already closed".to_string())),
None => Err(RestartError::Send("stream already closed".to_string())),
}
}
async fn recv(&mut self) -> Result<Option<AcquireRestartLockResponse>, RestartError> {
self.inbound
.message()
.await
.map_err(|status| RestartError::Recv(status.to_string()))
}
async fn close_send(&mut self) -> Result<(), RestartError> {
self.tx = None;
Ok(())
}
}
struct LockState {
token: String,
expires_at: Option<prost_types::Timestamp>,
}
pub struct Lock {
state: Arc<StdMutex<LockState>>,
lost_rx: watch::Receiver<Option<Arc<RestartError>>>,
release_tx: StdMutex<Option<oneshot::Sender<()>>>,
done_rx: StdMutex<Option<oneshot::Receiver<()>>>,
released: Arc<AtomicBool>,
}
impl std::fmt::Debug for Lock {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Lock")
.field("token", &self.token())
.field("expires_at", &self.expires_at())
.field("released", &self.released.load(Ordering::SeqCst))
.finish()
}
}
impl Lock {
pub fn token(&self) -> String {
self.state.lock().unwrap_or_else(|e| e.into_inner()).token.clone()
}
pub fn expires_at(&self) -> Option<prost_types::Timestamp> {
self.state.lock().unwrap_or_else(|e| e.into_inner()).expires_at
}
pub fn lost(&self) -> watch::Receiver<Option<Arc<RestartError>>> {
self.lost_rx.clone()
}
pub async fn release(&self) -> Result<(), RestartError> {
if self.released.swap(true, Ordering::SeqCst) {
return Ok(());
}
if let Some(tx) = self.release_tx.lock().unwrap_or_else(|e| e.into_inner()).take() {
let _ = tx.send(());
}
let done_rx = self.done_rx.lock().unwrap_or_else(|e| e.into_inner()).take();
if let Some(done_rx) = done_rx {
let _ = done_rx.await;
}
Ok(())
}
}
fn validate_ttl(ttl: Duration) -> Result<i32, RestartError> {
if ttl.is_zero() {
return Err(RestartError::InvalidTtl(ttl));
}
let secs = ttl.as_secs();
let secs = if secs == 0 { 1 } else { secs };
i32::try_from(secs).map_err(|_| RestartError::InvalidTtl(ttl))
}
pub async fn acquire_lock(
mut client: SecretsServiceClient<Channel>,
namespace: impl Into<String>,
service: impl Into<String>,
ttl: Duration,
) -> Result<Lock, RestartError> {
let ttl_secs = validate_ttl(ttl)?;
let namespace = namespace.into();
let service = service.into();
let (tx, rx) = mpsc::channel(4);
let response = client
.acquire_restart_lock(ReceiverStream::new(rx))
.await
.map_err(|status| RestartError::OpenStream(status.to_string()))?;
let stream = RealLockStream {
tx: Some(tx),
inbound: response.into_inner(),
};
acquire_lock_with(stream, namespace, service, ttl_secs).await
}
async fn acquire_lock_with<S: LockStream + Send + 'static>(
mut stream: S,
namespace: String,
service: String,
ttl_secs: i32,
) -> Result<Lock, RestartError> {
stream
.send(AcquireRestartLockRequest {
namespace,
service,
ttl_seconds: ttl_secs,
heartbeat: false,
})
.await?;
let (token, expires_at) = loop {
match stream.recv().await? {
Some(resp) if resp.message_type == MessageType::Acquired as i32 => {
break (resp.token, resp.expires_at);
}
Some(_) => continue,
None => return Err(RestartError::ClosedBeforeAcquired),
}
};
let state = Arc::new(StdMutex::new(LockState { token, expires_at }));
let (lost_tx, lost_rx) = watch::channel(None);
let (release_tx, release_rx) = oneshot::channel();
let (done_tx, done_rx) = oneshot::channel();
let released = Arc::new(AtomicBool::new(false));
tokio::spawn(run_lock_loop(
stream,
ttl_secs,
state.clone(),
lost_tx,
release_rx,
done_tx,
released.clone(),
));
Ok(Lock {
state,
lost_rx,
release_tx: StdMutex::new(Some(release_tx)),
done_rx: StdMutex::new(Some(done_rx)),
released,
})
}
async fn run_lock_loop<S: LockStream>(
mut stream: S,
ttl_secs: i32,
state: Arc<StdMutex<LockState>>,
lost_tx: watch::Sender<Option<Arc<RestartError>>>,
mut release_rx: oneshot::Receiver<()>,
done_tx: oneshot::Sender<()>,
released: Arc<AtomicBool>,
) {
let interval = heartbeat_interval(ttl_secs);
let mut ticker = tokio::time::interval_at(tokio::time::Instant::now() + interval, interval);
loop {
tokio::select! {
biased;
_ = &mut release_rx => {
let _ = stream.close_send().await;
break;
}
_ = ticker.tick() => {
if let Err(e) = stream.send(AcquireRestartLockRequest {
namespace: String::new(),
service: String::new(),
ttl_seconds: 0,
heartbeat: true,
}).await {
report_lost(&lost_tx, &released, e);
break;
}
}
recv_result = stream.recv() => {
match recv_result {
Ok(Some(resp)) => {
if resp.message_type == MessageType::TtlExtended as i32 {
let mut guard = state.lock().unwrap_or_else(|e| e.into_inner());
guard.expires_at = resp.expires_at;
}
}
Ok(None) => {
report_lost(&lost_tx, &released, RestartError::LockLost("stream closed by server".to_string()));
break;
}
Err(e) => {
report_lost(&lost_tx, &released, e);
break;
}
}
}
}
}
let _ = done_tx.send(());
}
fn heartbeat_interval(ttl_secs: i32) -> Duration {
let ttl = Duration::from_secs(ttl_secs.max(1) as u64);
let interval = ttl / 4;
if interval.is_zero() {
Duration::from_millis(1)
} else {
interval
}
}
fn report_lost(
lost_tx: &watch::Sender<Option<Arc<RestartError>>>,
released: &AtomicBool,
err: RestartError,
) {
if released.load(Ordering::SeqCst) {
return;
}
let _ = lost_tx.send(Some(Arc::new(err)));
}
trait WatchStream: Send {
fn recv(
&mut self,
) -> impl Future<Output = Result<Option<WatchServiceBundleResponse>, RestartError>> + Send;
}
struct RealWatchStream {
inbound: tonic::codec::Streaming<WatchServiceBundleResponse>,
}
impl WatchStream for RealWatchStream {
async fn recv(&mut self) -> Result<Option<WatchServiceBundleResponse>, RestartError> {
self.inbound
.message()
.await
.map_err(|status| RestartError::Recv(status.to_string()))
}
}
pub fn watch_bundle(
client: SecretsServiceClient<Channel>,
namespace: impl Into<String>,
service: impl Into<String>,
cancel: CancellationToken,
) -> mpsc::Receiver<()> {
let (changes_tx, changes_rx) = mpsc::channel(1);
let namespace = namespace.into();
let service = service.into();
tokio::spawn(async move {
watch_loop(changes_tx, cancel, move || {
let mut client = client.clone();
let namespace = namespace.clone();
let service = service.clone();
async move {
let response = client
.watch_service_bundle(WatchServiceBundleRequest { namespace, service })
.await
.map_err(|status| RestartError::WatchOpen(status.to_string()))?;
Ok::<_, RestartError>(RealWatchStream {
inbound: response.into_inner(),
})
}
})
.await;
});
changes_rx
}
async fn watch_loop<S, F, Fut>(changes: mpsc::Sender<()>, cancel: CancellationToken, mut open: F)
where
S: WatchStream + Send + 'static,
F: FnMut() -> Fut + Send,
Fut: Future<Output = Result<S, RestartError>> + Send,
{
let mut backoff = WATCH_BACKOFF_MIN;
loop {
if cancel.is_cancelled() {
return;
}
let mut stream = match open().await {
Ok(stream) => stream,
Err(_) => {
if !sleep_or_cancelled(backoff, &cancel).await {
return;
}
backoff = next_backoff(backoff);
continue;
}
};
loop {
tokio::select! {
() = cancel.cancelled() => return,
recv_result = stream.recv() => {
match recv_result {
Ok(Some(resp)) => {
backoff = WATCH_BACKOFF_MIN;
if resp.event_type == EventType::Changed as i32 {
let _ = changes.try_send(());
}
}
_ => break, }
}
}
}
if cancel.is_cancelled() {
return;
}
if !sleep_or_cancelled(backoff, &cancel).await {
return;
}
backoff = next_backoff(backoff);
}
}
fn next_backoff(current: Duration) -> Duration {
let doubled = current.saturating_mul(2);
if doubled > WATCH_BACKOFF_MAX {
WATCH_BACKOFF_MAX
} else {
doubled
}
}
async fn sleep_or_cancelled(d: Duration, cancel: &CancellationToken) -> bool {
tokio::select! {
() = cancel.cancelled() => false,
() = tokio::time::sleep(d) => true,
}
}
pub async fn wait_for_restart(
client: SecretsServiceClient<Channel>,
namespace: impl Into<String>,
service: impl Into<String>,
ttl: Duration,
debounce: Duration,
cancel: CancellationToken,
) -> Result<Lock, RestartError> {
let namespace = namespace.into();
let service = service.into();
let mut changes = watch_bundle(client.clone(), namespace.clone(), service.clone(), cancel.clone());
if changes.recv().await.is_none() {
return Err(RestartError::WatchStopped);
}
if !debounce.is_zero() {
let sleep = tokio::time::sleep(debounce);
tokio::pin!(sleep);
loop {
tokio::select! {
maybe_change = changes.recv() => {
match maybe_change {
Some(()) => sleep.as_mut().reset(tokio::time::Instant::now() + debounce),
None => return Err(RestartError::WatchStopped),
}
}
() = &mut sleep => break,
() = cancel.cancelled() => return Err(RestartError::Cancelled),
}
}
}
acquire_lock(client, namespace, service, ttl).await
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::VecDeque;
use std::sync::atomic::AtomicUsize;
use tonic::transport::Endpoint;
enum LockRecvResult {
Resp(AcquireRestartLockResponse),
Err(RestartError),
}
struct FakeLockStream {
rx: mpsc::UnboundedReceiver<LockRecvResult>,
sent: Arc<StdMutex<Vec<AcquireRestartLockRequest>>>,
closed: Arc<AtomicBool>,
}
impl LockStream for FakeLockStream {
async fn send(&mut self, req: AcquireRestartLockRequest) -> Result<(), RestartError> {
if self.closed.load(Ordering::SeqCst) {
return Err(RestartError::Send("send on closed stream".to_string()));
}
self.sent.lock().unwrap().push(req);
Ok(())
}
async fn recv(&mut self) -> Result<Option<AcquireRestartLockResponse>, RestartError> {
if self.closed.load(Ordering::SeqCst) {
return Err(RestartError::Recv("stream closed".to_string()));
}
match self.rx.recv().await {
Some(LockRecvResult::Resp(r)) => Ok(Some(r)),
Some(LockRecvResult::Err(e)) => Err(e),
None => Err(RestartError::Recv("stream closed".to_string())),
}
}
async fn close_send(&mut self) -> Result<(), RestartError> {
self.closed.store(true, Ordering::SeqCst);
Ok(())
}
}
struct FakeLockHandle {
tx: mpsc::UnboundedSender<LockRecvResult>,
sent: Arc<StdMutex<Vec<AcquireRestartLockRequest>>>,
#[allow(dead_code)]
closed: Arc<AtomicBool>,
}
impl FakeLockHandle {
fn new() -> (FakeLockStream, Self) {
let (tx, rx) = mpsc::unbounded_channel();
let sent = Arc::new(StdMutex::new(Vec::new()));
let closed = Arc::new(AtomicBool::new(false));
(
FakeLockStream {
rx,
sent: sent.clone(),
closed: closed.clone(),
},
Self { tx, sent, closed },
)
}
fn push_resp(&self, resp: AcquireRestartLockResponse) {
let _ = self.tx.send(LockRecvResult::Resp(resp));
}
fn push_err(&self, err: RestartError) {
let _ = self.tx.send(LockRecvResult::Err(err));
}
fn sent_count(&self) -> usize {
self.sent.lock().unwrap().len()
}
}
fn acquired_resp(token: &str, expires_at_secs_from_now: i64) -> AcquireRestartLockResponse {
AcquireRestartLockResponse {
message_type: MessageType::Acquired as i32,
position: 0,
token: token.to_string(),
expires_at: Some(future_timestamp(expires_at_secs_from_now)),
}
}
fn queue_position_resp(position: i32) -> AcquireRestartLockResponse {
AcquireRestartLockResponse {
message_type: MessageType::QueuePosition as i32,
position,
token: String::new(),
expires_at: None,
}
}
fn ttl_extended_resp(expires_at_secs_from_now: i64) -> AcquireRestartLockResponse {
AcquireRestartLockResponse {
message_type: MessageType::TtlExtended as i32,
position: 0,
token: String::new(),
expires_at: Some(future_timestamp(expires_at_secs_from_now)),
}
}
fn future_timestamp(secs_from_now: i64) -> prost_types::Timestamp {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap();
prost_types::Timestamp {
seconds: now.as_secs() as i64 + secs_from_now,
nanos: 0,
}
}
#[tokio::test]
async fn acquire_lock_rejects_ttl_zero_before_opening_any_stream() {
let channel = Endpoint::from_static("http://127.0.0.1:1").connect_lazy();
let client = SecretsServiceClient::new(channel);
let err = acquire_lock(client, "ns", "svc", Duration::ZERO)
.await
.unwrap_err();
assert!(
matches!(err, RestartError::InvalidTtl(_)),
"expected InvalidTtl, got {err:?}"
);
assert!(err.to_string().contains("ttl must be > 0"));
}
#[test]
fn validate_ttl_rejects_zero() {
let err = validate_ttl(Duration::ZERO).unwrap_err();
assert!(matches!(err, RestartError::InvalidTtl(_)));
}
#[test]
fn validate_ttl_floors_sub_second_ttl_to_one_second() {
assert_eq!(validate_ttl(Duration::from_millis(500)).unwrap(), 1);
}
#[test]
fn validate_ttl_truncates_like_go_int32_division() {
assert_eq!(validate_ttl(Duration::from_millis(4999)).unwrap(), 4);
}
#[tokio::test]
async fn acquire_lock_queue_position_then_acquired() {
let (stream, handle) = FakeLockHandle::new();
handle.push_resp(queue_position_resp(2));
handle.push_resp(queue_position_resp(1));
handle.push_resp(acquired_resp("tok-1", 30));
let lock = acquire_lock_with(stream, "ns".to_string(), "svc".to_string(), 4)
.await
.expect("acquire_lock_with");
assert_eq!(lock.token(), "tok-1");
assert_eq!(handle.sent_count(), 1, "expected exactly 1 initial send before any heartbeat");
lock.release().await.unwrap();
}
#[tokio::test]
async fn acquire_lock_stream_error_before_acquired_surfaces_as_err() {
let (stream, handle) = FakeLockHandle::new();
handle.push_err(RestartError::Recv("boom".to_string()));
let result = tokio::time::timeout(
Duration::from_secs(5),
acquire_lock_with(stream, "ns".to_string(), "svc".to_string(), 4),
)
.await
.expect("acquire_lock_with should not hang");
let err = result.unwrap_err();
assert!(err.to_string().contains("boom"));
}
#[tokio::test]
async fn acquire_lock_stream_closed_before_acquired_is_a_clear_error() {
let (stream, handle) = FakeLockHandle::new();
drop(handle);
let err = acquire_lock_with(stream, "ns".to_string(), "svc".to_string(), 4)
.await
.unwrap_err();
assert!(err.to_string().contains("stream closed"));
}
#[tokio::test]
async fn heartbeat_interval_is_ttl_over_four() {
let (stream, handle) = FakeLockHandle::new();
handle.push_resp(acquired_resp("tok-1", 4));
let lock = acquire_lock_with(stream, "ns".to_string(), "svc".to_string(), 1)
.await
.expect("acquire_lock_with");
tokio::time::sleep(Duration::from_millis(650)).await;
let sent = handle.sent_count();
assert!(sent >= 3, "sent_count={sent}, want >= 3 (1 initial + >=2 heartbeats)");
lock.release().await.unwrap();
}
#[tokio::test]
async fn ttl_extended_updates_expires_at() {
let (stream, handle) = FakeLockHandle::new();
handle.push_resp(acquired_resp("tok-1", 4));
let lock = acquire_lock_with(stream, "ns".to_string(), "svc".to_string(), 4)
.await
.expect("acquire_lock_with");
let extended = future_timestamp(60);
handle.push_resp(ttl_extended_resp(60));
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
loop {
if lock.expires_at().as_ref().map(|t| t.seconds) == Some(extended.seconds) {
break;
}
if tokio::time::Instant::now() > deadline {
panic!(
"expires_at never reflected TTL_EXTENDED; last = {:?}, want seconds = {}",
lock.expires_at(),
extended.seconds
);
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
lock.release().await.unwrap();
}
#[tokio::test]
async fn release_is_idempotent() {
let (stream, handle) = FakeLockHandle::new();
handle.push_resp(acquired_resp("tok-1", 30));
let lock = acquire_lock_with(stream, "ns".to_string(), "svc".to_string(), 4)
.await
.expect("acquire_lock_with");
lock.release().await.expect("first release");
lock.release().await.expect("second release should be a no-op");
}
#[tokio::test]
async fn lost_reported_on_unexpected_stream_error() {
let (stream, handle) = FakeLockHandle::new();
handle.push_resp(acquired_resp("tok-1", 4));
let lock = acquire_lock_with(stream, "ns".to_string(), "svc".to_string(), 4)
.await
.expect("acquire_lock_with");
handle.push_err(RestartError::Recv("lock lost".to_string()));
let mut lost = lock.lost();
tokio::time::timeout(Duration::from_secs(2), lost.changed())
.await
.expect("Lost() never signaled after stream error")
.expect("watch sender dropped unexpectedly");
let got = lost.borrow().clone();
assert!(got.is_some(), "expected Some(err) on Lost()");
}
#[tokio::test]
async fn release_does_not_report_lost() {
let (stream, handle) = FakeLockHandle::new();
handle.push_resp(acquired_resp("tok-1", 30));
let lock = acquire_lock_with(stream, "ns".to_string(), "svc".to_string(), 4)
.await
.expect("acquire_lock_with");
lock.release().await.expect("release");
assert!(
lock.lost().borrow().is_none(),
"Lost() recorded a value after a clean release: {:?}",
lock.lost().borrow()
);
}
#[tokio::test]
async fn lost_is_none_immediately_after_acquire() {
let (stream, handle) = FakeLockHandle::new();
handle.push_resp(acquired_resp("tok-1", 30));
let lock = acquire_lock_with(stream, "ns".to_string(), "svc".to_string(), 4)
.await
.expect("acquire_lock_with");
assert!(lock.lost().borrow().is_none());
lock.release().await.unwrap();
}
struct FakeWatchStream {
rx: mpsc::UnboundedReceiver<Result<WatchServiceBundleResponse, RestartError>>,
recv_count: Arc<AtomicUsize>,
}
impl WatchStream for FakeWatchStream {
async fn recv(&mut self) -> Result<Option<WatchServiceBundleResponse>, RestartError> {
self.recv_count.fetch_add(1, Ordering::SeqCst);
match self.rx.recv().await {
Some(Ok(r)) => Ok(Some(r)),
Some(Err(e)) => Err(e),
None => Err(RestartError::Recv("stream closed".to_string())),
}
}
}
fn changed_resp() -> WatchServiceBundleResponse {
WatchServiceBundleResponse {
event_type: EventType::Changed as i32,
namespace: "ns".to_string(),
service: "svc".to_string(),
}
}
#[tokio::test]
async fn watch_loop_coalesces_rapid_changes() {
let (tx, rx) = mpsc::unbounded_channel();
let recv_count = Arc::new(AtomicUsize::new(0));
for _ in 0..3 {
tx.send(Ok(changed_resp())).unwrap();
}
let stream = FakeWatchStream {
rx,
recv_count: recv_count.clone(),
};
let (changes_tx, mut changes_rx) = mpsc::channel(1);
let cancel = CancellationToken::new();
let cancel_for_loop = cancel.clone();
let streams = Arc::new(tokio::sync::Mutex::new(VecDeque::from(vec![stream])));
let handle = tokio::spawn(async move {
watch_loop(changes_tx, cancel_for_loop, move || {
let streams = streams.clone();
async move {
streams
.lock()
.await
.pop_front()
.ok_or_else(|| RestartError::WatchOpen("no more fake streams".to_string()))
}
})
.await;
});
let deadline = tokio::time::Instant::now() + Duration::from_secs(2);
while recv_count.load(Ordering::SeqCst) < 4 {
if tokio::time::Instant::now() > deadline {
panic!(
"watch_loop only issued {} Recv calls, want >= 4",
recv_count.load(Ordering::SeqCst)
);
}
tokio::time::sleep(Duration::from_millis(5)).await;
}
assert_eq!(changes_rx.len(), 1, "expected exactly 1 coalesced signal after 3 rapid changes");
changes_rx.recv().await.expect("one coalesced change");
assert!(
changes_rx.try_recv().is_err(),
"expected no further pending signal after draining the coalesced one"
);
cancel.cancel();
let _ = handle.await;
}
#[tokio::test]
async fn watch_loop_reconnects_after_stream_error() {
let (first_tx, first_rx) = mpsc::unbounded_channel();
let first_recv_count = Arc::new(AtomicUsize::new(0));
first_tx.send(Err(RestartError::Recv("boom".to_string()))).unwrap();
let first = FakeWatchStream {
rx: first_rx,
recv_count: first_recv_count,
};
let (second_tx, second_rx) = mpsc::unbounded_channel();
let second_recv_count = Arc::new(AtomicUsize::new(0));
second_tx.send(Ok(changed_resp())).unwrap();
let second = FakeWatchStream {
rx: second_rx,
recv_count: second_recv_count,
};
let (changes_tx, mut changes_rx) = mpsc::channel(1);
let cancel = CancellationToken::new();
let cancel_for_loop = cancel.clone();
let streams = Arc::new(tokio::sync::Mutex::new(VecDeque::from(vec![first, second])));
let handle = tokio::spawn(async move {
watch_loop(changes_tx, cancel_for_loop, move || {
let streams = streams.clone();
async move {
streams
.lock()
.await
.pop_front()
.ok_or_else(|| RestartError::WatchOpen("no more fake streams".to_string()))
}
})
.await;
});
tokio::time::timeout(Duration::from_secs(4), changes_rx.recv())
.await
.expect("expected watch_loop to reconnect and eventually deliver a change")
.expect("changes channel closed unexpectedly");
cancel.cancel();
let _ = handle.await;
}
#[test]
fn next_backoff_doubles_and_caps() {
assert_eq!(next_backoff(Duration::from_secs(1)), Duration::from_secs(2));
assert_eq!(next_backoff(Duration::from_secs(16)), Duration::from_secs(30));
assert_eq!(next_backoff(Duration::from_secs(30)), Duration::from_secs(30));
}
}