use std::net::{IpAddr, Ipv4Addr};
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use miden_protocol::block::BlockNumber;
use tokio::sync::{Semaphore, watch};
use tokio_stream::StreamExt;
use super::*;
struct TestSubscription {
ban_list: Arc<IpBanList>,
semaphore: Arc<Semaphore>,
fetch_count: Arc<AtomicUsize>,
fail_at: Option<BlockNumber>,
}
impl Default for TestSubscription {
fn default() -> Self {
Self {
ban_list: Arc::new(IpBanList::default()),
semaphore: Arc::new(Semaphore::new(1)),
fetch_count: Arc::new(AtomicUsize::new(0)),
fail_at: None,
}
}
}
impl TestSubscription {
fn failing_at(block: BlockNumber) -> Self {
Self { fail_at: Some(block), ..Self::default() }
}
fn fetch_count(&self) -> Arc<AtomicUsize> {
Arc::clone(&self.fetch_count)
}
fn stream(
&self,
from: BlockNumber,
chain_tip: watch::Receiver<BlockNumber>,
) -> tonic::Result<SubscriptionStream> {
self.stream_for_ip(None, from, chain_tip)
}
fn stream_for_ip(
&self,
client_ip: Option<IpAddr>,
from: BlockNumber,
chain_tip: watch::Receiver<BlockNumber>,
) -> tonic::Result<SubscriptionStream> {
let fetch_count = Arc::clone(&self.fetch_count);
let fail_at = self.fail_at;
SubscriptionStream::create(
from,
client_ip,
Arc::clone(&self.ban_list),
Arc::clone(&self.semaphore),
chain_tip,
move |block| {
fetch_count.fetch_add(1, Ordering::Relaxed);
let result = if Some(block) == fail_at {
Ok(None)
} else {
Ok(Some(block.as_u32().to_be_bytes().to_vec()))
};
std::future::ready(result)
},
)
}
}
#[tokio::test]
async fn stream_waiting_for_tip_returns_server_shutdown_when_tip_sender_closes() {
let (tip_tx, tip_rx) = watch::channel(BlockNumber::GENESIS);
let source = TestSubscription::default();
let mut stream = source
.stream(BlockNumber::from(1u32), tip_rx)
.expect("subscription start should be valid");
drop(tip_tx);
let item = tokio::time::timeout(Duration::from_secs(5), stream.next())
.await
.expect("stream must yield promptly")
.expect("stream must not end without an item");
assert_stream_status(item, tonic::Code::Unavailable);
wait_for_subscription_exit(&source).await;
}
#[tokio::test]
async fn stream_yields_requested_block_once_tip_reaches_it() {
let (_tip_tx, tip_rx) = watch::channel(BlockNumber::from(1u32));
let mut stream = TestSubscription::default()
.stream(BlockNumber::from(1u32), tip_rx)
.expect("subscription start should be valid");
let item = tokio::time::timeout(Duration::from_secs(5), stream.next())
.await
.expect("stream must yield promptly")
.expect("stream must not end without an item")
.expect("stream event must be ok");
assert_eq!(item.block, BlockNumber::from(1u32));
assert_eq!(item.tip, BlockNumber::from(1u32));
assert_eq!(item.data, 1u32.to_be_bytes().to_vec());
}
#[tokio::test]
async fn missing_data_returns_internal_and_reports_eos() {
let (_tip_tx, tip_rx) = watch::channel(BlockNumber::from(1u32));
let source = TestSubscription::failing_at(BlockNumber::from(1u32));
let mut stream = source
.stream(BlockNumber::from(1u32), tip_rx)
.expect("subscription start should be valid");
let item = tokio::time::timeout(Duration::from_secs(5), stream.next())
.await
.expect("stream must yield promptly")
.expect("stream must not end without an item");
assert_stream_status(item, tonic::Code::Internal);
wait_for_subscription_exit(&source).await;
}
#[tokio::test]
async fn stream_waiting_for_tip_exits_when_receiver_is_dropped() {
let (_tip_tx, tip_rx) = watch::channel(BlockNumber::GENESIS);
let source = TestSubscription::default();
let stream = source
.stream(BlockNumber::from(1u32), tip_rx)
.expect("subscription start should be valid");
drop(stream);
wait_for_subscription_exit(&source).await;
}
#[tokio::test]
async fn shutdown_while_send_is_pending_reports_server_shutdown() {
let (tip_tx, tip_rx) =
watch::channel(BlockNumber::from((SUBSCRIBER_CHANNEL_CAPACITY + 1) as u32));
let source = TestSubscription::default();
let fetch_count = source.fetch_count();
let _stream = source
.stream(BlockNumber::GENESIS, tip_rx)
.expect("subscription start should be valid");
wait_for_fetch_count(&fetch_count, SUBSCRIBER_CHANNEL_CAPACITY + 1).await;
drop(tip_tx);
wait_for_subscription_exit(&source).await;
}
#[tokio::test]
async fn slow_subscriber_is_banned() {
let (tip_tx, tip_rx) = watch::channel(BlockNumber::GENESIS);
let source = TestSubscription::default();
let client_ip = IpAddr::V4(Ipv4Addr::LOCALHOST);
let mut stream = source
.stream_for_ip(Some(client_ip), BlockNumber::GENESIS, tip_rx)
.expect("subscription start should be valid");
let first_item = tokio::time::timeout(Duration::from_secs(5), stream.next())
.await
.expect("stream must yield promptly")
.expect("stream must not end without an item")
.expect("stream event must be ok");
assert_eq!(first_item.block, BlockNumber::GENESIS);
tip_tx
.send(BlockNumber::from(SubscriberLagTracker::MAX_RUNNING_GAP + 2))
.expect("chain tip receiver should be open");
let item = tokio::time::timeout(Duration::from_secs(5), stream.next())
.await
.expect("stream must yield promptly")
.expect("stream must not end without an item");
assert_stream_status(item, tonic::Code::ResourceExhausted);
assert!(source.ban_list.banned_until(client_ip).is_some());
wait_for_subscription_exit(&source).await;
let (_tip_tx, tip_rx) = watch::channel(BlockNumber::GENESIS);
assert_subscription_start_err(
&source,
Some(client_ip),
BlockNumber::GENESIS,
tip_rx,
tonic::Code::ResourceExhausted,
);
}
#[tokio::test]
async fn subscription_start_future_gap_is_enforced() {
assert_subscription_start_ok(0, 10).await;
assert_subscription_start_ok(10, 10).await;
assert_subscription_start_ok(10 + MAX_FUTURE_GAP_IN_SUBSCRIPTIONS, 10).await;
let source = TestSubscription::default();
let (_tip_tx, tip_rx) = watch::channel(BlockNumber::from(10u32));
assert_subscription_start_err(
&source,
None,
BlockNumber::from(10 + MAX_FUTURE_GAP_IN_SUBSCRIPTIONS + 1),
tip_rx,
tonic::Code::OutOfRange,
);
assert_subscription_start_ok(u32::MAX, u32::MAX - 10).await;
}
async fn wait_for_subscription_exit(source: &TestSubscription) {
tokio::time::timeout(Duration::from_secs(5), async {
while source.semaphore.available_permits() == 0 {
tokio::task::yield_now().await;
}
})
.await
.expect("stream task must release subscription permit promptly");
}
async fn wait_for_fetch_count(fetch_count: &AtomicUsize, expected: usize) {
tokio::time::timeout(Duration::from_secs(5), async {
while fetch_count.load(Ordering::Relaxed) < expected {
tokio::task::yield_now().await;
}
})
.await
.expect("stream must fetch events promptly");
}
fn assert_stream_status(item: tonic::Result<StreamItem>, code: tonic::Code) {
let Err(err) = item else {
panic!("stream item must be an error");
};
assert_eq!(err.code(), code);
}
async fn assert_subscription_start_ok(block_from: u32, chain_tip: u32) {
let (tip_tx, tip_rx) = watch::channel(BlockNumber::from(chain_tip));
let source = TestSubscription::default();
let stream = source.stream(BlockNumber::from(block_from), tip_rx);
assert!(stream.is_ok());
drop(stream);
drop(tip_tx);
wait_for_subscription_exit(&source).await;
}
fn assert_subscription_start_err(
source: &TestSubscription,
client_ip: Option<IpAddr>,
block_from: BlockNumber,
chain_tip: watch::Receiver<BlockNumber>,
code: tonic::Code,
) {
match source.stream_for_ip(client_ip, block_from, chain_tip) {
Ok(_) => panic!("subscription start should be rejected"),
Err(err) => assert_eq!(err.code(), code),
}
}
fn run(gaps: &[u32]) -> bool {
let mut lag_tracker = SubscriberLagTracker::default();
for &gap in gaps {
if !lag_tracker.record_and_check(BlockNumber::GENESIS, BlockNumber::from(gap)) {
return false;
}
}
true
}
#[test]
fn starting_above_max_growth_is_ok() {
assert!(run(&[SubscriberLagTracker::MAX_RUNNING_GAP * 2]));
}
#[test]
fn accumulated_growth_limit_is_enforced() {
assert!(run(&[0, SubscriberLagTracker::MAX_RUNNING_GAP]));
assert!(!run(&[0, SubscriberLagTracker::MAX_RUNNING_GAP + 1]));
let step = SubscriberLagTracker::MAX_RUNNING_GAP / 4;
let gaps: Vec<u32> = (1..=6).map(|i| i * step).collect();
assert!(!run(&gaps));
}
#[test]
fn recovery_reduces_and_allows_fresh_accumulation() {
let near_limit = SubscriberLagTracker::MAX_RUNNING_GAP - 1;
assert!(run(&[near_limit, 1, near_limit]));
}
#[test]
fn token_improvement_does_not_prevent_disconnection() {
let gaps: Vec<u32> = (0u32..SubscriberLagTracker::MAX_RUNNING_GAP + 10)
.flat_map(|i| [50 + i, 49 + i])
.collect();
assert!(!run(&gaps));
}