mod chunks;
use chunks::*;
mod skip;
use skip::*;
mod step_by;
use step_by::*;
mod zip;
use tracing::trace;
use zip::*;
use crate::array::rdma::private::Sealed;
use crate::array::LamellarArray;
use crate::memregion::Dist;
use crate::warnings::RuntimeWarning;
use crate::LamellarTask;
use futures_util::{Future, Stream};
use pin_project::pin_project;
use std::pin::Pin;
use std::task::{Context, Poll};
pub(crate) mod private {
use crate::array::rdma::private::LamellarRdmaGet;
use crate::memregion::Dist;
use crate::{LamellarArray, LamellarEnv};
use std::pin::Pin;
use std::task::{Context, Poll};
pub trait OneSidedIteratorInner {
type Item: Send;
type ElemType: Dist + 'static;
type Array: LamellarRdmaGet<Self::ElemType>
+ LamellarArray<Self::ElemType>
+ LamellarEnv
+ Send;
fn init(&mut self);
fn next(&mut self) -> Option<Self::Item>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>>;
fn advance_index(&mut self, count: usize);
fn advance_index_pin(self: Pin<&mut Self>, count: usize);
fn array(&self) -> Self::Array;
fn item_size(&self) -> usize {
std::mem::size_of::<Self::Item>()
}
}
}
pub trait OneSidedIterator: private::OneSidedIteratorInner {
fn chunks(self, chunk_size: usize) -> Chunks<Self>
where
Self: Sized + Send,
{
Chunks::new(self, chunk_size)
}
fn skip(self, count: usize) -> Skip<Self>
where
Self: Sized + Send,
{
Skip::new(self, count)
}
fn step_by(self, step_size: usize) -> StepBy<Self>
where
Self: Sized + Send,
{
StepBy::new(self, step_size)
}
fn zip<I>(self, iter: I) -> Zip<Self, I>
where
Self: Sized + Send,
I: OneSidedIterator + Sized + Send,
{
Zip::new(self, iter)
}
fn into_iter(mut self) -> OneSidedIteratorIter<Self>
where
Self: Sized + Send,
{
RuntimeWarning::BlockingCall("into_iter", "into_stream()").print();
self.init();
OneSidedIteratorIter { iter: self }
}
fn into_stream(mut self) -> OneSidedStream<Self>
where
Self: Sized + Send,
{
self.init();
OneSidedStream { iter: self }
}
}
pub struct OneSidedIteratorIter<I> {
pub(crate) iter: I,
}
impl<I> Iterator for OneSidedIteratorIter<I>
where
I: OneSidedIterator,
{
type Item = <I as private::OneSidedIteratorInner>::Item;
fn next(&mut self) -> Option<Self::Item> {
self.iter.next()
}
}
#[pin_project]
pub struct OneSidedIter<T: Dist + 'static, A: LamellarArray<T>> {
array: A,
buf_size: usize,
index: usize,
buf_index: usize,
#[pin]
state: State<T>,
}
#[pin_project(project = StateProj)]
pub(crate) enum State<T> {
SinglePending(#[pin] LamellarTask<T>),
BufferedPending(#[pin] LamellarTask<Vec<T>>),
Buffered(Vec<T>),
Finished,
}
impl<T: Dist + 'static, A: LamellarArray<T>> OneSidedIter<T, A> {
pub(crate) fn new(array: &A, buf_size: usize) -> OneSidedIter<T, A> {
let iter = OneSidedIter {
array: array.clone(),
buf_size,
index: 0,
buf_index: 0,
state: State::Finished,
};
iter
}
}
impl<T: Dist + 'static + Clone + Send, A: LamellarArray<T> + Send> OneSidedIterator
for OneSidedIter<T, A>
{
}
impl<T: Dist + 'static + Clone + Send, A: LamellarArray<T> + Send> private::OneSidedIteratorInner
for OneSidedIter<T, A>
{
type ElemType = T;
type Item = T;
type Array = A;
fn init(&mut self) {
if self.buf_size == 1 {
let req = unsafe { self.array.get(self.index, Sealed) };
trace!("one sided iter init single get launched");
self.state = State::SinglePending(req.spawn());
} else {
trace!("one sided iter init buffered get launched");
let req = unsafe { self.array.get_buffer(self.index, self.buf_size, Sealed) };
self.state = State::BufferedPending(req.spawn());
}
}
fn next(&mut self) -> Option<Self::Item> {
let mut cur_state = State::Finished;
std::mem::swap(&mut self.state, &mut cur_state);
match cur_state {
State::SinglePending(req) => {
let data = req.block();
self.index += 1;
if self.index < self.array.len() {
trace!("one sided iter next single get launched");
let req = unsafe { self.array.get(self.index, Sealed) };
self.state = State::SinglePending(req.spawn());
} else {
trace!("one sided iter next finished");
self.state = State::Finished;
}
Some(data)
}
State::BufferedPending(req) => {
trace!("one sided iter buffered pending blocking");
let data = req.block();
let data_usize_slice = unsafe {
std::slice::from_raw_parts(data.as_ptr() as *const usize, data.len())
};
trace!(
"one sided iter buffered pending got data {:?}",
data_usize_slice
);
let val = data[0];
self.state = State::Buffered(data);
self.index += 1;
self.buf_index += 1;
Some(val)
}
State::Buffered(data) => {
let data_usize_slice = unsafe {
std::slice::from_raw_parts(data.as_ptr() as *const usize, data.len())
};
trace!(
"one sided iter buffered next: index: {} buf_index: {} data: {:?}",
self.index,
self.buf_index,
data_usize_slice
);
let val = data[self.buf_index];
self.index += 1;
self.buf_index += 1;
if self.index < self.array.len() {
if self.buf_index == self.buf_size {
self.buf_index = 0;
if self.index + self.buf_size < self.array.len() {
let req = unsafe {
self.array
.get_buffer(self.index, self.buf_size, Sealed)
.spawn()
};
self.state = State::BufferedPending(req);
} else {
let req = unsafe {
self.array
.get_buffer(self.index, self.array.len() - self.index, Sealed)
.spawn()
};
self.state = State::BufferedPending(req);
};
} else {
self.state = State::Buffered(data);
}
} else {
trace!("one sided iter buffered set finished");
self.state = State::Finished;
}
Some(val)
}
State::Finished => None,
}
}
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let mut this = self.project();
let res = match this.state.as_mut().project() {
StateProj::SinglePending(req) => match req.poll(cx) {
Poll::Ready(data) => {
let val = data;
*this.index += 1;
if *this.index < this.array.len() {
let req = unsafe { this.array.get(*this.index, Sealed).spawn() };
*this.state = State::SinglePending(req);
} else {
*this.state = State::Finished;
}
Some(val)
}
Poll::Pending => {
return Poll::Pending;
}
},
StateProj::BufferedPending(req) => match req.poll(cx) {
Poll::Ready(data) => {
let val = data[0];
*this.state = State::Buffered(data);
*this.index += 1;
*this.buf_index += 1;
Some(val)
}
Poll::Pending => {
return Poll::Pending;
}
},
StateProj::Buffered(data) => {
if *this.index < this.array.len() {
if *this.buf_index == *this.buf_size {
*this.buf_index = 0;
let req = if *this.index + *this.buf_size < this.array.len() {
unsafe {
this.array
.get_buffer(*this.index, *this.buf_size, Sealed)
.spawn()
}
} else {
unsafe {
this.array
.get_buffer(*this.index, this.array.len() - *this.index, Sealed)
.spawn()
}
};
*this.state = State::BufferedPending(req);
return Poll::Pending;
}
let val = data[*this.buf_index];
*this.index += 1;
*this.buf_index += 1;
Some(val)
} else {
*this.state = State::Finished;
None
}
}
StateProj::Finished => None,
};
Poll::Ready(res)
}
fn advance_index(&mut self, count: usize) {
let this = Pin::new(self);
this.advance_index_pin(count);
}
fn advance_index_pin(mut self: Pin<&mut Self>, count: usize) {
self.index += count;
if self.buf_size == 1 {
if self.index < self.array.len() {
let req = unsafe { self.array.get(self.index, Sealed).spawn() };
self.state = State::SinglePending(req);
} else {
self.state = State::Finished;
}
} else {
self.buf_index += count;
if self.index >= self.array.len() {
self.buf_index = 0;
self.state = State::Finished;
} else if self.buf_index == self.buf_size {
self.buf_index = 0;
if self.index + self.buf_size < self.array.len() {
let req = unsafe {
self.array
.get_buffer(self.index, self.buf_size, Sealed)
.spawn()
};
self.state = State::BufferedPending(req);
} else {
let req = unsafe {
self.array
.get_buffer(self.index, self.array.len() - self.index, Sealed)
.spawn()
};
self.state = State::BufferedPending(req);
}
}
}
}
fn array(&self) -> Self::Array {
self.array.clone()
}
fn item_size(&self) -> usize {
std::mem::size_of::<T>()
}
}
#[pin_project]
pub struct OneSidedStream<I> {
#[pin]
pub(crate) iter: I,
}
impl<I> Stream for OneSidedStream<I>
where
I: OneSidedIterator,
{
type Item = <I as private::OneSidedIteratorInner>::Item;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let this = self.project();
let res = this.iter.poll_next(cx);
match res {
Poll::Ready(Some(res)) => {
Poll::Ready(Some(res))
}
Poll::Ready(None) => {
Poll::Ready(None)
}
Poll::Pending => {
Poll::Pending
}
}
}
}