use crate::array::iterator::one_sided_iterator::*;
use crate::array::LamellarArrayRequest;
use crate::memregion::OneSidedMemoryRegion;
use std::collections::VecDeque;
use std::ops::Deref;
use async_trait::async_trait;
use pin_project::pin_project;
#[pin_project]
pub struct Buffered<I>
where
I: OneSidedIterator + Send,
{
#[pin]
iter: I,
index: usize,
buf_index: usize,
buf_size: usize,
reqs: VecDeque<Option<(usize, ArrayRdmaHandle, OneSidedMemoryRegion<u8>)>>,
state: BufferedState,
}
enum BufferedState {
Ready,
Pending,
Finished,
}
impl<I> Buffered<I>
where
I: OneSidedIterator + Send,
{
pub(crate) fn new(iter: I, buf_size: usize) -> Buffered<I> {
let mut buf = Buffered {
iter,
index: 0,
buf_index: 0,
buf_size: buf_size,
reqs: VecDeque::new(),
state: BufferedState::Pending,
};
for _ in 0..buf.buf_size {
buf.initiate_buffer();
}
buf
}
fn initiate_buffer(&mut self) {
let array = self.iter.array();
let array_bytes = array.len() * std::mem::size_of::<<Self as OneSidedIterator>::ElemType>();
let size = std::cmp::min(self.iter.item_size(), array_bytes - self.buf_index);
if size > 0 {
let mem_region = array.team().alloc_one_sided_mem_region(size);
if let Some(req) = self.iter.buffered_next(mem_region.clone()) {
self.reqs.push_back(Some((self.buf_index, req, mem_region)));
self.buf_index += size;
} else {
self.reqs.push_back(None);
}
}
}
fn wait_on_buffer(&mut self, size: usize) -> Option<OneSidedMemoryRegion<u8>> {
let (index, req, mem_region) =
if let Some((index, req, mem_region)) = self.reqs.pop_front().unwrap() {
(index, req, mem_region)
} else {
return None;
};
assert_eq!(mem_region.len(), size);
assert_eq!(index, self.index);
req.wait();
Some(mem_region)
}
}
pub struct BufferedItem<U> {
item: U,
_mem_region: OneSidedMemoryRegion<u8>,
}
impl<U> Deref for BufferedItem<U> {
type Target = U;
fn deref(&self) -> &Self::Target {
&self.item
}
}
impl<I> OneSidedIterator for Buffered<I>
where
I: OneSidedIterator + Send,
{
type ElemType = I::ElemType;
type Item = BufferedItem<I::Item>;
type Array = I::Array;
fn next(&mut self) -> Option<Self::Item> {
let array = self.array();
let array_bytes = array.len() * std::mem::size_of::<<Self as OneSidedIterator>::ElemType>();
if self.index < array_bytes {
let size = std::cmp::min(self.iter.item_size(), array_bytes - self.index);
let mem_region = self.wait_on_buffer(size)?;
self.index += size;
self.initiate_buffer();
Some(BufferedItem {
item: self.iter.from_mem_region(mem_region.clone())?,
_mem_region: mem_region,
})
} else {
None
}
}
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Ready(None)
}
fn advance_index(&mut self, count: usize) {
self.iter.advance_index(count);
}
fn array(&self) -> Self::Array {
self.iter.array()
}
fn item_size(&self) -> usize {
self.iter.item_size()
}
fn buffered_next(&mut self, mem_region: OneSidedMemoryRegion<u8>) -> Option<ArrayRdmaHandle> {
self.iter.buffered_next(mem_region)
}
fn from_mem_region(&self, mem_region: OneSidedMemoryRegion<u8>) -> Option<Self::Item> {
Some(BufferedItem {
item: self.iter.from_mem_region(mem_region.clone())?,
_mem_region: mem_region,
})
}
}