use std::sync::Arc;
use std::time::Duration;
use tokio::sync::{mpsc, Semaphore};
use tokio_util::sync::CancellationToken;
use crate::clock::Clock;
use crate::endpoint::Endpoint;
use crate::poller::{HttpClient, PollBatch, PollOutcome, PollResult};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum SchedulerRunError {
ReceiverDropped,
}
pub(crate) struct SchedulerRunHandle {
pub(crate) batches: mpsc::Receiver<PollBatch>,
#[allow(dead_code)]
pub(crate) task: tokio::task::JoinHandle<Result<(), SchedulerRunError>>,
}
pub struct PollScheduler<C: Clock> {
clock: C,
client: HttpClient,
refresh_interval: Duration,
max_concurrent: usize,
}
impl<C: Clock + Clone + Send + Sync + 'static> PollScheduler<C> {
#[must_use]
pub fn new(
clock: C,
client: HttpClient,
refresh_interval: Duration,
max_concurrent: usize,
) -> Self {
Self {
clock,
client,
refresh_interval,
max_concurrent,
}
}
pub fn run(
self,
endpoints: Vec<Endpoint>,
cancel: CancellationToken,
refresh_rx: mpsc::Receiver<()>,
) -> mpsc::Receiver<PollBatch> {
self.run_observed(endpoints, cancel, refresh_rx).batches
}
pub(crate) fn run_observed(
self,
endpoints: Vec<Endpoint>,
cancel: CancellationToken,
refresh_rx: mpsc::Receiver<()>,
) -> SchedulerRunHandle {
let (tx, rx) = mpsc::channel::<PollBatch>(4);
let task =
tokio::spawn(async move { self.poll_loop(endpoints, tx, cancel, refresh_rx).await });
SchedulerRunHandle { batches: rx, task }
}
async fn poll_loop(
self,
endpoints: Vec<Endpoint>,
tx: mpsc::Sender<PollBatch>,
cancel: CancellationToken,
mut refresh_rx: mpsc::Receiver<()>,
) -> Result<(), SchedulerRunError> {
if endpoints.is_empty() {
return Ok(());
}
let semaphore = Arc::new(Semaphore::new(self.max_concurrent));
let mut generation: u64 = 0;
let mut refresh_open = true;
let mut interval = tokio::time::interval(self.refresh_interval);
interval.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
generation = generation.saturating_add(1);
let batch = self
.poll_generation(&endpoints, &semaphore, generation)
.await;
if tokio::select! {
result = tx.send(batch) => result.is_err(),
() = cancel.cancelled() => false,
} {
return Err(SchedulerRunError::ReceiverDropped);
}
interval.tick().await;
loop {
tokio::select! {
biased;
() = cancel.cancelled() => break,
msg = refresh_rx.recv(), if refresh_open => {
match msg {
Some(()) => {
generation = generation.saturating_add(1);
let batch = self
.poll_generation(&endpoints, &semaphore, generation)
.await;
if tokio::select! {
result = tx.send(batch) => result.is_err(),
() = cancel.cancelled() => false,
} {
return Err(SchedulerRunError::ReceiverDropped);
}
}
None => {
refresh_open = false;
}
}
}
_ = interval.tick() => {
generation = generation.saturating_add(1);
let batch = self
.poll_generation(&endpoints, &semaphore, generation)
.await;
if tokio::select! {
result = tx.send(batch) => result.is_err(),
() = cancel.cancelled() => false,
} {
return Err(SchedulerRunError::ReceiverDropped);
}
}
}
}
Ok(())
}
async fn poll_generation(
&self,
endpoints: &[Endpoint],
semaphore: &Arc<Semaphore>,
generation: u64,
) -> PollBatch {
let started_at = self.clock.now();
let mut handles: Vec<(Endpoint, tokio::task::JoinHandle<PollResult>)> =
Vec::with_capacity(endpoints.len());
for endpoint in endpoints {
let client = self.client.clone();
let sem = Arc::clone(semaphore);
let ep = endpoint.clone();
let clock = self.clock.clone();
let handle = tokio::spawn(async move {
let _permit = sem.acquire().await.expect("semaphore should not be closed");
client.poll(&ep, &clock).await
});
handles.push((endpoint.clone(), handle));
}
let mut results = Vec::with_capacity(handles.len());
for (endpoint, handle) in handles {
match handle.await {
Ok(result) => results.push(result),
Err(_) => {
results.push(PollResult {
system_id: endpoint.id.clone(),
endpoint,
outcome: PollOutcome::Cancelled,
latency: Duration::ZERO,
});
}
}
}
PollBatch {
generation,
started_at,
completed_at: self.clock.now(),
results,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::clock::FakeClock;
use crate::endpoint::Endpoint;
use crate::poller::PollOutcome;
use gregg_protocol::test_support::LinuxSnapshotBuilder;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
fn refresh_channel() -> (mpsc::Sender<()>, mpsc::Receiver<()>) {
mpsc::channel(4)
}
async fn valid_snapshot_server() -> String {
let snap = LinuxSnapshotBuilder::default().build();
let body = serde_json::to_string(&snap).unwrap();
mock_server(body.into_bytes(), "200 OK").await
}
async fn mock_server(body: Vec<u8>, status: &str) -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let status = status.to_string();
tokio::spawn(async move {
loop {
let Ok((mut stream, _)) = listener.accept().await else {
break;
};
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let request = String::from_utf8_lossy(&buf[..total]);
let response_status = if request
.lines()
.next()
.is_some_and(|line| line.contains("/v2/"))
{
"404 Not Found"
} else {
&status
};
let header = format!(
"HTTP/1.1 {response_status}\r\nContent-Length: {}\r\n\r\n",
body.len()
);
stream.write_all(header.as_bytes()).await.unwrap();
stream.write_all(&body).await.unwrap();
}
});
format!("http://127.0.0.1:{}", addr.port())
}
fn endpoint_for_url(url: &str) -> Endpoint {
let stripped = url.strip_prefix("http://").unwrap();
let (host, port_str) = stripped.rsplit_once(':').unwrap();
Endpoint {
id: format!("{host}:{port_str}"),
host: host.to_string(),
port: port_str.parse().unwrap(),
name: None,
}
}
#[tokio::test]
async fn scheduler_produces_batches_with_increasing_generations() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let mut clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock.clone(), client, Duration::from_millis(10), 4);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let batch1 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(batch1.generation, 1);
clock.advance(Duration::from_millis(20));
let batch2 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(batch2.generation, 2);
cancel.cancel();
}
#[tokio::test]
async fn concurrency_never_exceeds_bound() {
let max_concurrent = 2;
let concurrent_count = Arc::new(AtomicUsize::new(0));
let peak_concurrent = Arc::new(AtomicUsize::new(0));
let mut endpoints = Vec::new();
for _ in 0..5 {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let cc = Arc::clone(&concurrent_count);
let pc = Arc::clone(&peak_concurrent);
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let current = cc.fetch_add(1, Ordering::SeqCst) + 1;
pc.fetch_max(current, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(50)).await;
cc.fetch_sub(1, Ordering::SeqCst);
let snap = LinuxSnapshotBuilder::default().build();
let body = serde_json::to_string(&snap).unwrap();
let header = format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len());
stream.write_all(header.as_bytes()).await.unwrap();
stream.write_all(body.as_bytes()).await.unwrap();
});
endpoints.push(Endpoint {
id: format!("ep-{}", addr.port()),
host: "127.0.0.1".into(),
port: addr.port(),
name: None,
});
}
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler =
PollScheduler::new(clock, client, Duration::from_millis(10), max_concurrent);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(endpoints, cancel.clone(), refresh_rx);
let _ = tokio::time::timeout(Duration::from_secs(5), rx.recv()).await;
cancel.cancel();
let peak = peak_concurrent.load(Ordering::SeqCst);
assert!(
peak <= max_concurrent,
"peak concurrent {peak} exceeded max {max_concurrent}"
);
}
#[tokio::test]
async fn cancellation_stops_scheduler() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock, client, Duration::from_millis(10), 4);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let batch = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap();
assert!(batch.is_some());
cancel.cancel();
tokio::time::sleep(Duration::from_millis(50)).await;
}
#[tokio::test]
async fn empty_endpoint_list() {
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock, client, Duration::from_millis(10), 4);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![], cancel.clone(), refresh_rx);
let result = tokio::time::timeout(Duration::from_millis(100), rx.recv()).await;
assert!(result.unwrap().is_none());
cancel.cancel();
}
#[tokio::test]
async fn single_endpoint_polls_repeatedly() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let mut clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock.clone(), client, Duration::from_millis(10), 4);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let mut generations = Vec::new();
for _ in 0..3 {
clock.advance(Duration::from_millis(20));
if let Some(batch) = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
{
generations.push(batch.generation);
}
}
assert_eq!(generations, vec![1, 2, 3]);
cancel.cancel();
}
#[tokio::test]
async fn overlap_skip_if_running() {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
tokio::time::sleep(Duration::from_millis(100)).await;
let snap = LinuxSnapshotBuilder::default().build();
let body = serde_json::to_string(&snap).unwrap();
let header = format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len());
stream.write_all(header.as_bytes()).await.unwrap();
stream.write_all(body.as_bytes()).await.unwrap();
});
let ep = Endpoint {
id: "slow-ep".into(),
host: "127.0.0.1".into(),
port: addr.port(),
name: None,
};
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let mut clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock.clone(), client, Duration::from_millis(20), 4);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let batch1 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(batch1.generation, 1);
clock.advance(Duration::from_millis(60));
clock.advance(Duration::from_millis(100));
let batch2 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(batch2.generation, 2);
cancel.cancel();
}
#[tokio::test]
async fn multiple_endpoints_all_polled() {
let url1 = valid_snapshot_server().await;
let url2 = valid_snapshot_server().await;
let ep1 = endpoint_for_url(&url1);
let ep2 = endpoint_for_url(&url2);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let mut clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock.clone(), client, Duration::from_millis(10), 4);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep1, ep2], cancel.clone(), refresh_rx);
clock.advance(Duration::from_millis(20));
let batch = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(batch.results.len(), 2);
cancel.cancel();
}
#[tokio::test]
async fn fleet_scaling_10_endpoints() {
fleet_scaling_test(10, 4).await;
}
#[tokio::test]
async fn fleet_scaling_50_endpoints() {
fleet_scaling_test(50, 4).await;
}
#[tokio::test]
async fn fleet_scaling_100_endpoints() {
fleet_scaling_test(100, 4).await;
}
async fn fleet_scaling_test(n: usize, max_concurrent: usize) {
let mut endpoints = Vec::new();
for _ in 0..n {
let url = valid_snapshot_server().await;
endpoints.push(endpoint_for_url(&url));
}
let client = HttpClient::new(Duration::from_secs(30));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler =
PollScheduler::new(clock, client, Duration::from_millis(10), max_concurrent);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(endpoints, cancel.clone(), refresh_rx);
let batch = tokio::time::timeout(Duration::from_secs(60), rx.recv())
.await
.expect("should receive batch within timeout")
.expect("channel should not be closed");
assert_eq!(
batch.results.len(),
n,
"should have one result per endpoint"
);
let online_count = batch
.results
.iter()
.filter(|r| matches!(r.outcome, PollOutcome::Online(_)))
.count();
assert_eq!(online_count, n, "all endpoints should be online");
cancel.cancel();
}
#[tokio::test]
async fn fleet_scaling_concurrency_bounded_at_scale() {
let n = 50;
let max_concurrent = 4;
let concurrent_count = Arc::new(AtomicUsize::new(0));
let peak_concurrent = Arc::new(AtomicUsize::new(0));
let mut endpoints = Vec::new();
for _ in 0..n {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let cc = Arc::clone(&concurrent_count);
let pc = Arc::clone(&peak_concurrent);
tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let current = cc.fetch_add(1, Ordering::SeqCst) + 1;
pc.fetch_max(current, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(20)).await;
cc.fetch_sub(1, Ordering::SeqCst);
let snap = LinuxSnapshotBuilder::default().build();
let body = serde_json::to_string(&snap).unwrap();
let header = format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len());
stream.write_all(header.as_bytes()).await.unwrap();
stream.write_all(body.as_bytes()).await.unwrap();
});
endpoints.push(Endpoint {
id: format!("ep-{}", addr.port()),
host: "127.0.0.1".into(),
port: addr.port(),
name: None,
});
}
let client = HttpClient::new(Duration::from_secs(30));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler =
PollScheduler::new(clock, client, Duration::from_millis(10), max_concurrent);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(endpoints, cancel.clone(), refresh_rx);
let batch = tokio::time::timeout(Duration::from_secs(60), rx.recv())
.await
.expect("should receive batch")
.expect("channel open");
assert_eq!(batch.results.len(), n);
cancel.cancel();
let peak = peak_concurrent.load(Ordering::SeqCst);
assert!(
peak <= max_concurrent,
"peak concurrent {peak} exceeded max {max_concurrent}"
);
}
async fn alternating_mock_server() -> String {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let snap = LinuxSnapshotBuilder::default().build();
let body = serde_json::to_string(&snap).unwrap();
let call_count = Arc::new(AtomicUsize::new(0));
tokio::spawn(async move {
loop {
let Ok((mut stream, _)) = listener.accept().await else {
break;
};
let mut buf = vec![0u8; 4096];
let mut total = 0;
loop {
let n = stream.read(&mut buf[total..]).await.unwrap();
total += n;
if buf[..total].windows(4).any(|w| w == b"\r\n\r\n") {
break;
}
}
let request = String::from_utf8_lossy(&buf[..total]);
if request
.lines()
.next()
.is_some_and(|line| line.contains("/v2/"))
{
let header = "HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\n\r\n";
stream.write_all(header.as_bytes()).await.unwrap();
continue;
}
let count = call_count.fetch_add(1, Ordering::SeqCst);
if count % 2 == 0 {
let header =
format!("HTTP/1.1 200 OK\r\nContent-Length: {}\r\n\r\n", body.len());
stream.write_all(header.as_bytes()).await.unwrap();
stream.write_all(body.as_bytes()).await.unwrap();
} else {
drop(stream);
}
}
});
format!("http://127.0.0.1:{}", addr.port())
}
#[tokio::test]
async fn alternating_online_offline_endpoint() {
let url = alternating_mock_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let clock = crate::clock::RealClock;
let mut online_count = 0;
let mut offline_count = 0;
for _ in 0..6 {
let result = client.poll(&ep, &clock).await;
match &result.outcome {
PollOutcome::Online(_) => online_count += 1,
_ => offline_count += 1,
}
}
assert!(online_count > 0, "should have at least one online result");
assert!(offline_count > 0, "should have at least one offline result");
}
#[tokio::test]
async fn clock_backward_adjustment_does_not_corrupt_scheduler() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let mut clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock.clone(), client, Duration::from_millis(10), 4);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
clock.advance(Duration::from_millis(20));
let batch1 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(batch1.generation, 1);
assert!(batch1.started_at <= batch1.completed_at);
clock.set(anchor.checked_sub(Duration::from_secs(3600)).unwrap());
clock.advance(Duration::from_millis(20));
let batch2 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(batch2.generation, 2, "generations must be monotonic");
clock.set(anchor + Duration::from_secs(7200));
clock.advance(Duration::from_millis(20));
let batch3 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
.unwrap();
assert_eq!(batch3.generation, 3, "generations must be monotonic");
cancel.cancel();
}
#[tokio::test]
async fn scheduler_handles_alternating_endpoint() {
let url = alternating_mock_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let mut clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock.clone(), client, Duration::from_millis(10), 4);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let mut online_results = 0;
let mut offline_results = 0;
for _ in 0..4 {
clock.advance(Duration::from_millis(20));
if let Some(batch) = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.unwrap()
{
for result in &batch.results {
match &result.outcome {
PollOutcome::Online(_) => online_results += 1,
_ => offline_results += 1,
}
}
}
}
assert!(online_results > 0, "should have at least one online result");
assert!(
offline_results > 0,
"should have at least one offline result"
);
cancel.cancel();
}
#[tokio::test]
async fn first_poll_happens_immediately_without_delay() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock, client, Duration::from_secs(3600), 4);
let cancel = CancellationToken::new();
let (_refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let batch = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("should receive first batch without delay")
.expect("channel should be open");
assert_eq!(batch.generation, 1);
cancel.cancel();
}
#[tokio::test]
async fn refresh_now_triggers_generation() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock, client, Duration::from_secs(3600), 4);
let cancel = CancellationToken::new();
let (refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let batch1 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("first batch")
.expect("channel open");
assert_eq!(batch1.generation, 1);
refresh_tx.send(()).await.unwrap();
let batch2 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("refresh batch")
.expect("channel open");
assert_eq!(batch2.generation, 2);
cancel.cancel();
}
#[tokio::test]
async fn panicked_task_produces_cancelled_result_for_endpoint() {
use crate::poller::PollResult;
let endpoint = Endpoint {
id: "panic-ep".into(),
host: "127.0.0.1".into(),
port: 1,
name: None,
};
let ep_clone = endpoint.clone();
let panic_handle = tokio::spawn(async move {
let id = ep_clone.id.clone();
drop(id);
panic!("test panic");
});
let mut results = Vec::new();
match panic_handle.await {
Ok(result) => results.push(result),
Err(_) => {
results.push(PollResult {
system_id: endpoint.id.clone(),
endpoint,
outcome: PollOutcome::Cancelled,
latency: Duration::ZERO,
});
}
}
assert_eq!(results.len(), 1);
assert_eq!(results[0].system_id, "panic-ep");
assert_eq!(results[0].outcome, PollOutcome::Cancelled);
}
#[tokio::test]
async fn one_refresh_signal_produces_one_generation() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock, client, Duration::from_secs(3600), 4);
let cancel = CancellationToken::new();
let (refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let batch1 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("first batch")
.expect("channel open");
assert_eq!(batch1.generation, 1);
refresh_tx.send(()).await.unwrap();
let batch2 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("refresh batch")
.expect("channel open");
assert_eq!(batch2.generation, 2);
let result = tokio::time::timeout(Duration::from_millis(200), rx.recv()).await;
assert!(
result.is_err() || result.unwrap().is_none(),
"should not receive a third batch from a single refresh signal"
);
cancel.cancel();
}
#[tokio::test]
async fn closed_refresh_channel_does_not_busy_loop() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock, client, Duration::from_millis(200), 4);
let cancel = CancellationToken::new();
let (refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let batch1 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("first batch")
.expect("channel open");
assert_eq!(batch1.generation, 1);
let _t1 = std::time::Instant::now();
drop(refresh_tx);
let batch2 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("second batch")
.expect("channel open");
assert_eq!(batch2.generation, 2);
let t2 = std::time::Instant::now();
let batch3 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("third batch")
.expect("channel open");
assert_eq!(batch3.generation, 3);
let t3 = std::time::Instant::now();
let elapsed = t3.saturating_duration_since(t2);
assert!(
elapsed >= Duration::from_millis(100),
"generations should be spaced by the interval, not busy-looping; elapsed = {elapsed:?}"
);
cancel.cancel();
}
#[tokio::test]
async fn manual_refresh_does_not_reset_periodic_cadence() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_secs(5));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock, client, Duration::from_millis(100), 4);
let cancel = CancellationToken::new();
let (refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let batch1 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("first batch")
.expect("channel open");
assert_eq!(batch1.generation, 1);
let t1 = std::time::Instant::now();
let batch2 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("second batch")
.expect("channel open");
assert_eq!(batch2.generation, 2);
let _t2 = std::time::Instant::now();
refresh_tx.send(()).await.unwrap();
let batch3 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("refresh batch")
.expect("channel open");
assert_eq!(batch3.generation, 3);
let _t3 = std::time::Instant::now();
let batch4 = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("fourth batch")
.expect("channel open");
assert_eq!(batch4.generation, 4);
let t4 = std::time::Instant::now();
let elapsed_from_start = t4.saturating_duration_since(t1);
assert!(
elapsed_from_start < Duration::from_secs(30),
"periodic cadence should not be reset by manual refresh; \
elapsed from start = {elapsed_from_start:?}"
);
cancel.cancel();
}
#[tokio::test]
async fn three_rapid_refresh_signals_produce_three_generations() {
let url = valid_snapshot_server().await;
let ep = endpoint_for_url(&url);
let client = HttpClient::new(Duration::from_millis(100));
let anchor = std::time::Instant::now();
let clock = FakeClock::new(anchor);
let scheduler = PollScheduler::new(clock, client, Duration::from_secs(3600), 4);
let cancel = CancellationToken::new();
let (refresh_tx, refresh_rx) = refresh_channel();
let mut rx = scheduler.run(vec![ep], cancel.clone(), refresh_rx);
let _ = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("first batch")
.expect("channel open");
for _ in 0..3 {
refresh_tx.send(()).await.unwrap();
}
let mut generations = Vec::new();
for _ in 0..3 {
let batch = tokio::time::timeout(Duration::from_secs(5), rx.recv())
.await
.expect("batch")
.expect("channel open");
generations.push(batch.generation);
}
assert_eq!(generations, vec![2, 3, 4]);
cancel.cancel();
}
}