use std::cell::OnceCell;
use std::collections::VecDeque;
use std::ops::Range;
use slab::Slab;
use crate::common::ext::aligned_vec::ACow;
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::{
DiskCache, DiskCacheRemote, block_aligned_fetch, to_block_range,
};
use crate::common::universal_io::{ReadPipeline, UioResult, UniversalIoError, UniversalRead, UserData};
#[cfg(target_os = "linux")]
pub(super) const REMOTE_READ_ALIGNMENT: usize = crate::common::universal_io::io_uring::KERNEL_PAGE_SIZE;
#[cfg(not(target_os = "linux"))]
pub(super) const REMOTE_READ_ALIGNMENT: usize = 1;
struct InFlightFetch<'file, R, U>
where
R: UniversalRead + 'static,
{
file: &'file DiskCache<R>,
blocks_range: Range<u32>,
reads: Vec<(U, Range<u64>)>,
}
enum Source {
Local {
range: Range<u64>,
is_sequential: bool,
},
Remote {
blocks_range: Range<u32>,
blocks_byte_range: Range<u64>,
},
}
fn pick_source<P>(local: &LocalState, range: Range<u64>) -> UioResult<Source>
where
P: AccessPattern,
{
if range.is_empty() {
return Ok(Source::Local {
range,
is_sequential: P::IS_SEQUENTIAL,
});
}
if range.end > local.mmap().len::<u8>()? {
return Err(UniversalIoError::OutOfBounds {
start: range.start,
end: range.end,
elements: (range.end - range.start) as usize,
});
}
let blocks_range = to_block_range(range.clone());
if local.contains(blocks_range.clone()) {
return Ok(Source::Local {
range,
is_sequential: P::IS_SEQUENTIAL,
});
}
let (blocks_range, blocks_byte_range) = block_aligned_fetch(range, local.mmap().len::<u8>()?)
.expect("non-empty range has a non-empty block range");
Ok(Source::Remote {
blocks_range,
blocks_byte_range,
})
}
unsafe fn read_local<R>(
file: &DiskCache<R>,
range: Range<u64>,
is_sequential: bool,
) -> UioResult<&[u8]>
where
R: DiskCacheRemote,
{
if range.is_empty() {
return Ok(&[]);
}
let local = file.state()?.local;
if is_sequential {
unsafe { local.read_mmap_bytes::<Sequential>(range) }
} else {
unsafe { local.read_mmap_bytes::<Random>(range) }
}
}
unsafe fn commit_and_read<'file, R, U>(
fetch: InFlightFetch<'file, R, U>,
bytes: &[u8],
results: &mut VecDeque<(U, &'file [u8])>,
) -> UioResult<()>
where
R: DiskCacheRemote,
U: UserData,
{
let InFlightFetch {
file,
blocks_range,
reads,
} = fetch;
let local = file.state()?.local;
unsafe {
local.write_mmap_bytes(bytes, blocks_range);
for (user_data, read_range) in reads {
let slice = local.read_mmap_bytes::<Random>(read_range)?;
results.push_back((user_data, slice));
}
}
Ok(())
}
type RemotePipeline<'file, R> = <R as UniversalRead>::ReadPipeline<'file, u64>;
pub struct DiskCachePipeline<'file, R, U>
where
R: UniversalRead + 'static,
U: UserData,
{
remote_pipeline: OnceCell<RemotePipeline<'file, R>>,
in_flight: Slab<InFlightFetch<'file, R, U>>,
results: VecDeque<(U, &'file [u8])>,
}
impl<'file, R, U> DiskCachePipeline<'file, R, U>
where
R: UniversalRead + 'file,
U: UserData,
{
fn get_or_init_remote_pipeline<'a>(
remote_pipeline: &'a mut OnceCell<RemotePipeline<'file, R>>,
) -> UioResult<&'a mut RemotePipeline<'file, R>> {
if remote_pipeline.get().is_none() {
let remote = R::ReadPipeline::new()?;
let _ = remote_pipeline.set(remote);
}
Ok(remote_pipeline.get_mut().expect("just initialized"))
}
#[cfg(test)]
pub(super) fn in_flight_fetches(&self) -> usize {
self.in_flight.len()
}
}
impl<'file, R, U> ReadPipeline<'file, U> for DiskCachePipeline<'file, R, U>
where
R: DiskCacheRemote + 'file,
U: UserData,
{
type File = DiskCache<R>;
fn new() -> UioResult<Self> {
Ok(Self {
remote_pipeline: OnceCell::new(),
in_flight: Slab::new(),
results: VecDeque::new(),
})
}
fn can_schedule(&mut self) -> bool {
self.results.is_empty()
&& 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: Range<u64>,
_align: usize,
) -> UioResult<()> {
let state = file.state()?;
match pick_source::<P>(state.local, range.clone())? {
Source::Local {
range,
is_sequential,
} => {
let bytes = unsafe { read_local::<R>(file, range, is_sequential)? };
self.results.push_back((user_data, bytes));
}
Source::Remote {
blocks_range,
blocks_byte_range,
} => {
if let Some((_, fetch)) = self.in_flight.iter_mut().find(|(_, inflight)| {
std::ptr::eq(inflight.file, file)
&& inflight.blocks_range.start <= blocks_range.start
&& blocks_range.end <= inflight.blocks_range.end
}) {
fetch.reads.push((user_data, range));
return Ok(());
}
let remote_pipeline = Self::get_or_init_remote_pipeline(&mut self.remote_pipeline)?;
let entry = self.in_flight.vacant_entry();
remote_pipeline.schedule::<P>(
entry.key() as u64,
state.remote,
blocks_byte_range,
REMOTE_READ_ALIGNMENT,
)?;
entry.insert(InFlightFetch {
file,
blocks_range,
reads: vec![(user_data, range)],
});
}
}
Ok(())
}
fn schedule_whole(
&mut self,
user_data: U,
file: &'file DiskCache<R>,
from: u64,
) -> UioResult<()>
where
Self::File: UniversalRead,
{
let state = file.state()?;
let eof = state.local.mmap().len::<u8>()?;
if from >= eof {
return Ok(());
}
self.schedule::<Sequential>(user_data, file, from..eof, 1)
}
fn wait(&mut self) -> UioResult<Option<(U, ACow<'file>)>> {
if let Some((user_data, slice)) = self.results.pop_front() {
return Ok(Some((user_data, ACow::Borrowed(slice))));
}
let Some(remote_pipeline) = self.remote_pipeline.get_mut() else {
return Ok(None);
};
let completion = match remote_pipeline.wait() {
Ok(completion) => completion,
Err(err) => {
self.in_flight.clear();
self.remote_pipeline.take();
return Err(err);
}
};
let Some((fetch_id, bytes)) = completion else {
return Ok(None);
};
let fetch = self
.in_flight
.try_remove(fetch_id as usize)
.expect("completed fetch has an in-flight entry");
unsafe { commit_and_read(fetch, &bytes, &mut self.results)? };
let (user_data, slice) = self
.results
.pop_front()
.expect("a completed fetch resolves at least one read");
Ok(Some((user_data, ACow::Borrowed(slice))))
}
}