use std::borrow::Cow;
use std::cell::OnceCell;
use std::ops::Range;
use crate::common::generic_consts::{AccessPattern, Random, Sequential};
use crate::common::universal_io::simple_disk_cache::local_state::LocalState;
use crate::common::universal_io::simple_disk_cache::{BLOCK_SIZE, DiskCache, to_block_range};
use crate::common::universal_io::traits::BorrowedReadPipeline;
use crate::common::universal_io::{
self, Item, OwnedReadPipeline, ReadRange, Result, UniversalIoError, UniversalRead,
UniversalReadFs, UserData,
};
struct RemoteMeta<File, U> {
file: File,
blocks_range: Range<u32>,
read_range: ReadRange,
user_data: U,
}
enum Source {
Local {
range: ReadRange,
is_sequential: bool,
},
Remote {
blocks_range: Range<u32>,
blocks_byte_range: ReadRange,
},
}
fn pick_source<P, T>(local: &LocalState, range: ReadRange) -> Result<Source>
where
P: AccessPattern,
T: bytemuck::Pod,
{
if range.length == 0 {
return Ok(Source::Local {
range,
is_sequential: P::IS_SEQUENTIAL,
});
}
let byte_range = range.into_byte_range::<T>();
if byte_range.end > local.mmap().len::<u8>()? {
return Err(UniversalIoError::OutOfBounds {
start: byte_range.start,
end: byte_range.end,
elements: range.length as usize,
});
}
let blocks_range = to_block_range(byte_range.clone());
if local.contains(blocks_range.clone()) {
return Ok(Source::Local {
range,
is_sequential: P::IS_SEQUENTIAL,
});
}
let byte_offset = blocks_range.start as usize * BLOCK_SIZE;
let fetch_length = blocks_range.len() * BLOCK_SIZE;
let max_length = local.mmap().len::<u8>()?.saturating_sub(byte_offset as u64);
let blocks_byte_range = ReadRange {
byte_offset: byte_offset as u64,
length: max_length.min(fetch_length as u64),
};
Ok(Source::Remote {
blocks_range,
blocks_byte_range,
})
}
unsafe fn read_local<R, T>(
file: &DiskCache<R>,
range: ReadRange,
is_sequential: bool,
) -> universal_io::Result<&[T]>
where
R: UniversalRead + Clone,
R::Fs: Clone + Send + Sync,
<R::Fs as UniversalReadFs>::OpenExtra: Clone + Send + Sync,
R::OwnedReadPipeline<u8, Range<u32>>: Send,
T: bytemuck::Pod,
{
if range.length == 0 {
return Ok(&[]);
}
let local = file.local_state()?;
if is_sequential {
unsafe { local.read_mmap_bytes::<Sequential, T>(range) }
} else {
unsafe { local.read_mmap_bytes::<Random, T>(range) }
}
}
unsafe fn commit_and_read<'a, R, T>(
file: &'a DiskCache<R>,
bytes: &[u8],
blocks_range: Range<u32>,
read_range: ReadRange,
) -> universal_io::Result<&'a [T]>
where
R: UniversalRead + Clone,
R::Fs: Clone + Send + Sync,
<R::Fs as UniversalReadFs>::OpenExtra: Clone + Send + Sync,
R::OwnedReadPipeline<u8, Range<u32>>: Send,
T: bytemuck::Pod,
{
let local = file.local_state()?;
unsafe {
local.write_mmap_bytes(bytes, blocks_range);
local.read_mmap_bytes::<Random, T>(read_range)
}
}
type BorrowedRemotePipeline<'file, R, U> =
<R as UniversalRead>::BorrowedReadPipeline<'file, u8, RemoteMeta<&'file DiskCache<R>, U>>;
pub struct DiskCachePipeline<'file, R, T, U>
where
R: UniversalRead,
T: bytemuck::Pod,
U: UserData,
{
remote_pipeline: OnceCell<BorrowedRemotePipeline<'file, R, U>>,
result: Option<(U, &'file [T])>,
}
impl<'file, R, T, U> DiskCachePipeline<'file, R, T, U>
where
R: UniversalRead + 'file,
T: bytemuck::Pod,
U: UserData,
{
fn get_or_init_remote_pipeline(
&mut self,
) -> universal_io::Result<&mut BorrowedRemotePipeline<'file, R, U>> {
if self.remote_pipeline.get().is_none() {
let remote = R::BorrowedReadPipeline::new()?;
let _ = self.remote_pipeline.set(remote);
}
Ok(self.remote_pipeline.get_mut().expect("just initialized"))
}
}
impl<'file, R, T, U> BorrowedReadPipeline<'file, T, U> for DiskCachePipeline<'file, R, T, U>
where
R: UniversalRead + Clone + 'file,
R::Fs: Clone + Send + Sync,
<R::Fs as UniversalReadFs>::OpenExtra: Clone + Send + Sync,
R::OwnedReadPipeline<u8, Range<u32>>: Send,
T: Item,
{
type File = DiskCache<R>;
fn new() -> universal_io::Result<Self> {
Ok(Self {
remote_pipeline: OnceCell::new(),
result: None,
})
}
fn can_schedule(&mut self) -> bool {
self.result.is_none()
&& self
.remote_pipeline
.get_mut()
.is_none_or(|remote| remote.can_schedule())
}
fn schedule<P: AccessPattern>(
&mut self,
user_data: U,
file: &'file DiskCache<R>,
range: ReadRange,
) -> universal_io::Result<()> {
match pick_source::<P, T>(file.local_state()?, range)? {
Source::Local {
range,
is_sequential,
} => {
let bytes = unsafe { read_local::<R, T>(file, range, is_sequential)? };
self.result = Some((user_data, bytes));
}
Source::Remote {
blocks_range,
blocks_byte_range,
} => {
let remote_meta = RemoteMeta {
file,
blocks_range,
read_range: range,
user_data,
};
let remote_pipeline = self.get_or_init_remote_pipeline()?;
remote_pipeline.schedule::<P>(remote_meta, file.remote()?, blocks_byte_range)?;
}
}
Ok(())
}
fn wait(&mut self) -> universal_io::Result<Option<(U, Cow<'file, [T]>)>> {
if let Some((user_data, slice)) = self.result.take() {
return Ok(Some((user_data, Cow::Borrowed(slice))));
}
let Some(remote_pipeline) = self.remote_pipeline.get_mut() else {
return Ok(None);
};
let Some((remote_meta, bytes)) = remote_pipeline.wait()? else {
return Ok(None);
};
let RemoteMeta {
file,
blocks_range,
read_range,
user_data,
} = remote_meta;
let items = unsafe { commit_and_read::<R, T>(file, &bytes, blocks_range, read_range)? };
Ok(Some((user_data, Cow::Borrowed(items))))
}
}
pub struct OwnedDiskCachePipeline<R, T, U>
where
R: UniversalRead,
T: bytemuck::Pod,
U: UserData,
{
file: DiskCache<R>,
remote_pipeline: OnceCell<R::OwnedReadPipeline<u8, RemoteMeta<(), U>>>,
ready: Option<(U, ReadRange, bool)>,
_phantom: std::marker::PhantomData<T>,
}
impl<R, T, U> OwnedDiskCachePipeline<R, T, U>
where
R: UniversalRead + Clone,
R::Fs: Clone + Send + Sync,
<R::Fs as UniversalReadFs>::OpenExtra: Clone + Send + Sync,
R::OwnedReadPipeline<u8, Range<u32>>: Send,
T: bytemuck::Pod,
U: UserData,
{
fn get_or_init_remote_pipeline(
&mut self,
) -> universal_io::Result<&mut R::OwnedReadPipeline<u8, RemoteMeta<(), U>>> {
if self.remote_pipeline.get().is_none() {
let remote = R::OwnedReadPipeline::new(self.file.remote()?.clone())?;
let _ = self.remote_pipeline.set(remote);
}
Ok(self.remote_pipeline.get_mut().expect("just initialized"))
}
}
impl<R, T, U> OwnedReadPipeline<T, U> for OwnedDiskCachePipeline<R, T, U>
where
R: UniversalRead + Clone,
R::Fs: Clone + Send + Sync,
<R::Fs as UniversalReadFs>::OpenExtra: Clone + Send + Sync,
R::OwnedReadPipeline<u8, Range<u32>>: Send,
T: bytemuck::Pod,
{
type File = DiskCache<R>;
fn new(file: Self::File) -> universal_io::Result<Self> {
Ok(Self {
file,
remote_pipeline: OnceCell::new(),
ready: None,
_phantom: std::marker::PhantomData,
})
}
fn can_schedule(&mut self) -> bool {
self.ready.is_none()
&& self
.remote_pipeline
.get_mut()
.is_none_or(|remote| remote.can_schedule())
}
fn schedule<P>(&mut self, user_data: U, range: ReadRange) -> universal_io::Result<()>
where
P: AccessPattern,
{
match pick_source::<P, T>(self.file.local_state()?, range)? {
Source::Local {
range,
is_sequential,
} => {
self.ready = Some((user_data, range, is_sequential));
}
Source::Remote {
blocks_range,
blocks_byte_range,
} => {
let remote_meta = RemoteMeta {
file: (),
blocks_range,
read_range: range,
user_data,
};
let remote_pipeline = self.get_or_init_remote_pipeline()?;
remote_pipeline.schedule::<P>(remote_meta, blocks_byte_range)?;
}
}
Ok(())
}
fn schedule_whole(&mut self, user_data: U) -> Result<()> {
let length = self.file.len::<T>()?;
self.schedule::<Sequential>(
user_data,
ReadRange {
byte_offset: 0,
length,
},
)
}
fn wait(&mut self) -> universal_io::Result<Option<(U, Cow<'_, [T]>)>> {
if let Some((user_data, range, is_sequential)) = self.ready.take() {
let items = unsafe { read_local::<R, T>(&self.file, range, is_sequential)? };
return Ok(Some((user_data, Cow::Borrowed(items))));
}
let Some(remote_pipeline) = self.remote_pipeline.get_mut() else {
return Ok(None);
};
let Some((remote_meta, bytes)) = remote_pipeline.wait()? else {
return Ok(None);
};
let RemoteMeta {
file: _,
blocks_range,
read_range,
user_data,
} = remote_meta;
let items =
unsafe { commit_and_read::<R, T>(&self.file, &bytes, blocks_range, read_range)? };
Ok(Some((user_data, Cow::Borrowed(items))))
}
}