use super::*;
use std::{
net::SocketAddr,
sync::{
atomic::{AtomicUsize, Ordering},
Arc,
},
time::Duration,
};
#[test]
fn test_limits_concurrency() {
let mut state = RehandshakeState::new_with_paused_time(Duration::from_secs(60)).unwrap();
for i in 0..10 {
state.push(SocketAddr::from(([127, 0, 0, 1], 4000 + i)));
}
state.adjust_post_refill();
let concurrent = Arc::new(AtomicUsize::new(0));
let max_concurrent = Arc::new(AtomicUsize::new(0));
let handle = state.runtime.handle().clone();
state.next_rehandshake_batch(10, |_addr| {
let concurrent = concurrent.clone();
let max_concurrent = max_concurrent.clone();
Some(handle.spawn(async move {
let current = concurrent.fetch_add(1, Ordering::Relaxed) + 1;
max_concurrent.fetch_max(current, Ordering::Relaxed);
tokio::time::sleep(Duration::from_secs(5)).await;
concurrent.fetch_sub(1, Ordering::Relaxed);
}))
});
assert_eq!(max_concurrent.load(Ordering::Relaxed), 2);
}
#[test]
fn test_limits_concurrency_for_fast_handshakes() {
let mut state = RehandshakeState::new_with_paused_time(Duration::from_secs(60)).unwrap();
for i in 0..10 {
state.push(SocketAddr::from(([127, 0, 0, 1], 4000 + i)));
}
state.adjust_post_refill();
let concurrent = Arc::new(AtomicUsize::new(0));
let max_concurrent = Arc::new(AtomicUsize::new(0));
let handle = state.runtime.handle().clone();
state.next_rehandshake_batch(10, |_addr| {
let concurrent = concurrent.clone();
let max_concurrent = max_concurrent.clone();
Some(handle.spawn(async move {
let current = concurrent.fetch_add(1, Ordering::Relaxed) + 1;
max_concurrent.fetch_max(current, Ordering::Relaxed);
tokio::time::sleep(Duration::from_millis(1)).await;
concurrent.fetch_sub(1, Ordering::Relaxed);
}))
});
assert_eq!(max_concurrent.load(Ordering::Relaxed), 1);
}
#[test]
fn test_waits_for_completion() {
let mut state = RehandshakeState::new_with_paused_time(Duration::from_secs(60)).unwrap();
for i in 0..5 {
state.push(SocketAddr::from(([127, 0, 0, 1], 4000 + i)));
}
state.adjust_post_refill();
let completed = Arc::new(AtomicUsize::new(0));
let handle = state.runtime.handle().clone();
state.next_rehandshake_batch(5, |_addr| {
let completed = completed.clone();
Some(handle.spawn(async move {
tokio::time::sleep(Duration::from_secs(3)).await;
completed.fetch_add(1, Ordering::Relaxed);
}))
});
assert_eq!(completed.load(Ordering::Relaxed), 5);
}
#[test]
fn test_respects_60_second_deadline() {
let mut state = RehandshakeState::new_with_paused_time(Duration::from_secs(120)).unwrap();
for i in 0..100 {
state.push(SocketAddr::from(([127, 0, 0, 1], 4000 + i)));
}
state.adjust_post_refill();
let scheduled = Arc::new(AtomicUsize::new(0));
let handle = state.runtime.handle().clone();
state.next_rehandshake_batch(100, |_addr| {
let scheduled = scheduled.clone();
scheduled.fetch_add(1, Ordering::Relaxed);
Some(handle.spawn(async move {
tokio::time::sleep(Duration::from_secs(10)).await;
}))
});
let count = scheduled.load(Ordering::Relaxed);
assert!(count <= 13, "scheduled {count} handshakes, expected <= 13");
assert!(!state.queue.is_empty());
}
#[test]
fn test_keeps_unscheduled_in_queue() {
let mut state = RehandshakeState::new_with_paused_time(Duration::from_secs(60)).unwrap();
for i in 0..20 {
state.push(SocketAddr::from(([127, 0, 0, 1], 4000 + i)));
}
state.adjust_post_refill();
let initial_count = state.queue.len();
let handle = state.runtime.handle().clone();
state.next_rehandshake_batch(20, |_addr| {
Some(handle.spawn(async move {
tokio::time::sleep(Duration::from_secs(30)).await;
}))
});
let scheduled = initial_count - state.queue.len();
assert!(
scheduled < initial_count,
"should not schedule all handshakes"
);
assert!(
!state.queue.is_empty(),
"should keep unscheduled items in queue"
);
}
#[test]
fn test_cancelled_task_does_not_panic() {
let mut state = RehandshakeState::new_with_paused_time(Duration::from_secs(60)).unwrap();
state.push(SocketAddr::from(([127, 0, 0, 1], 4000)));
let handle = state.runtime.handle().clone();
state.next_rehandshake_batch(1, |_addr| {
let h = handle.spawn(async {
tokio::time::sleep(Duration::from_secs(60)).await;
});
h.abort();
Some(h)
});
}
#[test]
fn test_tail_handshake_scheduling() {
let mut state = RehandshakeState::new_with_paused_time(Duration::from_secs(3600)).unwrap();
state.push(SocketAddr::from(([127, 0, 0, 1], 4000)));
let scheduled = Arc::new(AtomicUsize::new(0));
let handle = state.runtime.handle().clone();
for _ in 0..61 {
state.next_rehandshake_batch(1, |_addr| {
scheduled.fetch_add(1, Ordering::Relaxed);
Some(handle.spawn(async move {}))
});
state.runtime.block_on(async {
tokio::time::sleep(Duration::from_secs(60)).await;
});
}
assert!(
scheduled.load(Ordering::Relaxed) > 0,
"should schedule at least one tail handshake"
);
}