use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use aerospike_rt::Mutex;
use async_channel::{Receiver, Sender};
use crate::errors::Result;
use crate::query::{PartitionFilter, PartitionTracker};
use crate::Record;
pub struct RecordStream {
rs: Arc<Recordset>,
rx: std::pin::Pin<Box<Receiver<Result<Record>>>>,
}
#[derive(Debug)]
pub struct Recordset {
instances: AtomicUsize,
rx: Receiver<Result<Record>>,
tx: Sender<Result<Record>>,
active: AtomicBool,
task_id: AtomicU64,
pub(crate) tracker: Arc<Mutex<PartitionTracker>>,
}
impl Drop for Recordset {
fn drop(&mut self) {
self.close();
}
}
impl Recordset {
pub(crate) fn new(
rec_queue_size: usize,
max_records: u64,
nodes: usize,
tracker: Arc<Mutex<PartitionTracker>>,
) -> Self {
let task_id = rand::random::<u64>();
let capacity = if max_records > 0 {
rec_queue_size.min(max_records as usize)
} else {
rec_queue_size
};
let (tx, rx) = async_channel::bounded(capacity.max(1));
Recordset {
instances: AtomicUsize::new(nodes),
rx,
tx,
active: AtomicBool::new(true),
task_id: AtomicU64::new(task_id),
tracker,
}
}
pub fn close(&self) {
self.active.store(false, Ordering::Relaxed);
self.rx.close();
}
pub fn is_active(&self) -> bool {
self.active.load(Ordering::Relaxed)
}
pub(crate) fn set_instances(&self, count: usize) {
self.instances.store(count, Ordering::Relaxed);
}
pub(crate) fn reset_task_id(&self) {
let task_id = rand::random::<u64>();
self.task_id.store(task_id, Ordering::Relaxed);
}
pub(crate) async fn err(&self, e: crate::Error) {
let _ = self.tx.clone().send(Err(e)).await;
}
pub(crate) async fn push(&self, record: Result<Record>) -> Result<()> {
match record {
Err(crate::Error::StreamTerminatedError()) => Ok(()),
_ => match self.tx.send(record).await {
Ok(()) => Ok(()),
Err(_) => Err(crate::Error::StreamTerminatedError()),
},
}
}
pub(crate) fn task_id(&self) -> u64 {
self.task_id.load(Ordering::Relaxed)
}
pub(crate) fn signal_end(&self) {
if self.instances.fetch_sub(1, Ordering::Relaxed) == 1 {
self.close();
}
}
pub async fn partition_filter(&self) -> Option<PartitionFilter> {
if !self.is_active() {
return self.tracker.lock().await.extract_partition_filter();
}
None
}
#[cfg(feature = "sync")]
pub fn next_record(&self) -> Option<Result<Record>> {
self.rx.try_recv().ok()
}
pub fn into_stream(self: Arc<Self>) -> RecordStream {
let rx = Box::pin(self.rx.clone());
RecordStream { rs: self, rx }
}
}
#[cfg(feature = "sync")]
impl Iterator for &Recordset {
type Item = Result<Record>;
fn next(&mut self) -> Option<Result<Record>> {
self.rx.recv_blocking().ok()
}
}
impl futures::Stream for RecordStream {
type Item = Result<Record>;
fn poll_next(
self: std::pin::Pin<&mut Self>,
cx: &mut std::task::Context<'_>,
) -> std::task::Poll<Option<Self::Item>> {
self.get_mut().rx.as_mut().poll_next(cx)
}
}
impl AsRef<Recordset> for RecordStream {
fn as_ref(&self) -> &Recordset {
&self.rs
}
}
impl RecordStream {
pub async fn partition_filter(&self) -> Option<PartitionFilter> {
self.rs.partition_filter().await
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
use std::time::Duration;
use futures::executor::block_on;
use futures::StreamExt;
use crate::policy::QueryPolicy;
fn recordset(queue_size: usize) -> Arc<Recordset> {
let tracker = block_on(PartitionTracker::new(
&QueryPolicy::default(),
Arc::new(Mutex::new(PartitionFilter::all())),
Vec::new(),
))
.expect("tracker");
Arc::new(Recordset::new(
queue_size,
0,
1,
Arc::new(Mutex::new(tracker)),
))
}
fn recordset_with_max(queue_size: usize, max_records: u64) -> Arc<Recordset> {
let tracker = block_on(PartitionTracker::new(
&QueryPolicy::default(),
Arc::new(Mutex::new(PartitionFilter::all())),
Vec::new(),
))
.expect("tracker");
Arc::new(Recordset::new(
queue_size,
max_records,
1,
Arc::new(Mutex::new(tracker)),
))
}
fn record() -> Record {
Record::new(None, HashMap::new(), 0, 0)
}
#[test]
fn queue_is_capped_by_max_records() {
let rs = recordset_with_max(1024, 10);
assert_eq!(rs.rx.capacity(), Some(10));
}
#[test]
fn queue_keeps_its_size_when_max_records_is_larger() {
let rs = recordset_with_max(64, 10_000);
assert_eq!(rs.rx.capacity(), Some(64));
}
#[test]
fn queue_keeps_its_size_when_max_records_is_zero() {
let rs = recordset_with_max(64, 0);
assert_eq!(rs.rx.capacity(), Some(64));
}
#[test]
fn queue_never_degenerates_to_zero() {
assert_eq!(recordset_with_max(0, 0).rx.capacity(), Some(1));
assert_eq!(recordset_with_max(1024, 0).rx.capacity(), Some(1024));
assert_eq!(recordset_with_max(0, 10).rx.capacity(), Some(1));
}
#[cfg(feature = "sync")]
#[test]
fn blocking_iterator_drains_buffered_records_after_close() {
let rs = recordset(8);
for _ in 0..3 {
block_on(rs.push(Ok(record()))).unwrap();
}
rs.close();
let mut iter = &*rs;
assert!(iter.next().is_some());
assert!(iter.next().is_some());
assert!(iter.next().is_some());
assert!(iter.next().is_none());
assert!(iter.next().is_none());
}
#[cfg(feature = "sync")]
#[test]
fn blocking_iterator_ends_immediately_on_closed_empty_set() {
let rs = recordset(8);
rs.close();
assert!((&*rs).next().is_none());
}
#[cfg(feature = "sync")]
#[test]
fn parked_iterator_wakes_on_close() {
let rs = recordset(8);
let (done_tx, done_rx) = std::sync::mpsc::channel();
let consumer_rs = rs.clone();
std::thread::spawn(move || {
let item = (&*consumer_rs).next(); let _ = done_tx.send(item.is_none());
});
std::thread::sleep(Duration::from_millis(100));
rs.close();
let ended_clean = done_rx
.recv_timeout(Duration::from_secs(5))
.expect("parked iterator was not woken by close()");
assert!(ended_clean, "expected None after close on empty set");
}
#[test]
fn stream_parks_on_an_empty_queue_instead_of_waking_itself() {
use futures::Stream;
use std::task::{Context, Poll, Wake, Waker};
struct Counter(AtomicUsize);
impl Wake for Counter {
fn wake(self: Arc<Self>) {
self.0.fetch_add(1, Ordering::SeqCst);
}
}
let counter = Arc::new(Counter(AtomicUsize::new(0)));
let waker = Waker::from(counter.clone());
let mut cx = Context::from_waker(&waker);
let rs = recordset(4);
let mut stream = rs.clone().into_stream();
let mut stream = std::pin::Pin::new(&mut stream);
assert!(matches!(stream.as_mut().poll_next(&mut cx), Poll::Pending));
assert!(matches!(stream.as_mut().poll_next(&mut cx), Poll::Pending));
assert_eq!(counter.0.load(Ordering::SeqCst), 0, "a parked stream must not wake itself");
block_on(rs.push(Ok(record()))).expect("push");
assert!(counter.0.load(Ordering::SeqCst) >= 1, "push must wake the parked stream");
assert!(matches!(stream.as_mut().poll_next(&mut cx), Poll::Ready(Some(Ok(_)))));
block_on(rs.push(Ok(record()))).expect("push before close");
rs.close();
assert!(matches!(stream.as_mut().poll_next(&mut cx), Poll::Ready(Some(Ok(_)))));
assert!(matches!(stream.as_mut().poll_next(&mut cx), Poll::Ready(None)));
}
#[test]
fn push_fails_fast_after_close() {
let rs = recordset(8);
rs.close();
let err = block_on(rs.push(Ok(record()))).unwrap_err();
assert!(
matches!(err, crate::Error::StreamTerminatedError()),
"unexpected error: {0}",
err
);
}
#[test]
fn producer_blocked_on_full_queue_unblocks_on_close() {
let rs = recordset(1);
block_on(rs.push(Ok(record()))).unwrap();
let (done_tx, done_rx) = std::sync::mpsc::channel();
let producer_rs = rs.clone();
std::thread::spawn(move || {
let result = block_on(producer_rs.push(Ok(record())));
let _ = done_tx.send(result.is_err());
});
std::thread::sleep(Duration::from_millis(100));
rs.close();
let send_failed = done_rx
.recv_timeout(Duration::from_secs(5))
.expect("blocked producer was not unblocked by close()");
assert!(send_failed, "push into a closed recordset must fail");
}
#[test]
fn async_stream_ends_after_close_and_drain() {
let rs = recordset(8);
for _ in 0..2 {
block_on(rs.push(Ok(record()))).unwrap();
}
rs.close();
let mut stream = rs.into_stream();
assert!(block_on(stream.next()).is_some());
assert!(block_on(stream.next()).is_some());
assert!(block_on(stream.next()).is_none());
}
#[test]
#[ignore = "stress test; run explicitly with --ignored"]
fn final_record_survives_a_close_racing_the_poll() {
let trials: usize = std::env::var("RACE_TRIALS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(20_000);
let burners: usize = std::env::var("RACE_BURNERS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(0);
let stop = Arc::new(AtomicBool::new(false));
let burner_handles: Vec<_> = (0..burners)
.map(|_| {
let stop = stop.clone();
std::thread::spawn(move || {
while !stop.load(Ordering::Relaxed) {
std::hint::spin_loop();
}
})
})
.collect();
let mut lost = 0usize;
for _ in 0..trials {
let rs = recordset(8);
let consumer_rs = rs.clone();
let polling = Arc::new(AtomicBool::new(false));
let polling_signal = polling.clone();
let consumer = std::thread::spawn(move || {
let mut stream = consumer_rs.into_stream();
block_on(async move {
let mut seen = 0usize;
while stream.next().await.is_some() {
seen += 1;
if seen == 1 {
polling_signal.store(true, Ordering::Release);
}
}
seen
})
});
block_on(rs.push(Ok(record()))).expect("warm-up push");
while !polling.load(Ordering::Acquire) {
std::hint::spin_loop();
}
for _ in 0..(rand::random::<u32>() % 4096) {
std::hint::spin_loop();
}
block_on(rs.push(Ok(record()))).expect("final push before close");
rs.close();
if consumer.join().expect("consumer thread") != 2 {
lost += 1;
}
}
stop.store(true, Ordering::Relaxed);
for handle in burner_handles {
let _ = handle.join();
}
assert_eq!(
lost, 0,
"{lost} of {trials} trials ended the stream with a record still queued",
);
}
}