check-32bit-casts.sh flagged two casts added by the previous commits; both values are bounded (checked by gather_storage's first walk, and built from a usize fetch length), so they go through addr::saturating_usize. Co-Authored-By: Claude Opus 5.5 (1M context) <[email protected]>
598 lines
21 KiB
Rust
598 lines
21 KiB
Rust
//! Copying a selection out of a row-major buffer one contiguous run at a time.
|
|
//!
|
|
//! A selection's elements, in output order, fall into runs that are adjacent
|
|
//! in the source: a whole block along the last dimension, blocks that touch
|
|
//! (`stride == block`), and whole rows when the inner dimensions are selected
|
|
//! in full. Copying run by run turns a 256 x 256 hyperslab of a 1024-wide
|
|
//! dataset into 256 `memcpy`s of 1 KiB, where the old extractor recursed and
|
|
//! bounds-checked once per element.
|
|
|
|
#[cfg(not(feature = "std"))]
|
|
use alloc::{vec, vec::Vec};
|
|
|
|
use crate::data_read::NativeElement;
|
|
use crate::error::FormatError;
|
|
use crate::selection::Selection;
|
|
use crate::storage::{ExtentBytes, ExtentReq, Storage, raw_batches};
|
|
|
|
/// Row-major element strides of `dims` (the last dimension has stride 1).
|
|
fn strides(dims: &[u64]) -> Vec<u64> {
|
|
let mut s = vec![1u64; dims.len()];
|
|
for d in (0..dims.len().saturating_sub(1)).rev() {
|
|
s[d] = s[d + 1].wrapping_mul(dims[d + 1]);
|
|
}
|
|
s
|
|
}
|
|
|
|
/// Merges adjacent runs before handing them on.
|
|
struct Coalesce<F: FnMut(u64, u64)> {
|
|
start: u64,
|
|
len: u64,
|
|
emit: F,
|
|
}
|
|
|
|
impl<F: FnMut(u64, u64)> Coalesce<F> {
|
|
#[inline]
|
|
fn push(&mut self, start: u64, len: u64) {
|
|
if len == 0 {
|
|
return;
|
|
}
|
|
if self.len > 0 && self.start.wrapping_add(self.len) == start {
|
|
self.len += len;
|
|
return;
|
|
}
|
|
self.flush();
|
|
self.start = start;
|
|
self.len = len;
|
|
}
|
|
|
|
fn flush(&mut self) {
|
|
if self.len > 0 {
|
|
(self.emit)(self.start, self.len);
|
|
self.len = 0;
|
|
}
|
|
}
|
|
}
|
|
|
|
/// Call `emit(first_element, element_count)` for each run of a hyperslab's
|
|
/// elements that is contiguous in a row-major dataset of shape `dims`, in
|
|
/// the order the selection returns them. Adjacent runs are merged.
|
|
///
|
|
/// Coordinates at or past a dimension's extent are skipped, as the
|
|
/// element-wise extractor always did; callers that want them to be an error
|
|
/// validate the selection first. The four vectors must have `dims.len()`
|
|
/// entries.
|
|
pub(crate) fn hyperslab_runs(
|
|
dims: &[u64],
|
|
start: &[u64],
|
|
stride: &[u64],
|
|
count: &[u64],
|
|
block: &[u64],
|
|
emit: impl FnMut(u64, u64),
|
|
) {
|
|
let rank = dims.len();
|
|
let mut out = Coalesce {
|
|
start: 0,
|
|
len: 0,
|
|
emit,
|
|
};
|
|
if rank == 0 {
|
|
out.push(0, 1);
|
|
out.flush();
|
|
return;
|
|
}
|
|
if (0..rank).any(|d| count[d] == 0 || block[d] == 0) {
|
|
return;
|
|
}
|
|
let strides = strides(dims);
|
|
let last = rank - 1;
|
|
// Odometer over the outer dimensions: (block index, offset in block).
|
|
let mut ci = vec![0u64; last];
|
|
let mut bi = vec![0u64; last];
|
|
'outer: loop {
|
|
// Base offset of this row, or skip it if a coordinate is out of range.
|
|
let mut base = 0u64;
|
|
let mut in_range = true;
|
|
for d in 0..last {
|
|
let coord = start[d]
|
|
.saturating_add(ci[d].saturating_mul(stride[d]))
|
|
.saturating_add(bi[d]);
|
|
if coord >= dims[d] {
|
|
in_range = false;
|
|
break;
|
|
}
|
|
base = base.wrapping_add(coord.wrapping_mul(strides[d]));
|
|
}
|
|
if in_range && (stride[last] == block[last] || count[last] == 1) {
|
|
// Blocks that touch (the common unit-stride case: block 1,
|
|
// stride 1) are one range; don't split it into per-element runs.
|
|
let s = start[last];
|
|
let e = s
|
|
.saturating_add(count[last].saturating_mul(block[last]))
|
|
.min(dims[last]);
|
|
if s < e {
|
|
out.push(base.wrapping_add(s), e - s);
|
|
}
|
|
} else if in_range {
|
|
for c in 0..count[last] {
|
|
let s = start[last].saturating_add(c.saturating_mul(stride[last]));
|
|
if s >= dims[last] {
|
|
continue;
|
|
}
|
|
let e = s.saturating_add(block[last]).min(dims[last]);
|
|
out.push(base.wrapping_add(s), e - s);
|
|
}
|
|
}
|
|
// Advance the odometer, last outer dimension fastest.
|
|
let mut d = last;
|
|
loop {
|
|
if d == 0 {
|
|
break 'outer;
|
|
}
|
|
d -= 1;
|
|
bi[d] += 1;
|
|
if bi[d] < block[d] {
|
|
break;
|
|
}
|
|
bi[d] = 0;
|
|
ci[d] += 1;
|
|
if ci[d] < count[d] {
|
|
break;
|
|
}
|
|
ci[d] = 0;
|
|
}
|
|
}
|
|
out.flush();
|
|
}
|
|
|
|
/// The selected elements of `src` — a row-major dataset of shape `dims` and
|
|
/// `elem_size`-byte elements — copied into a fresh `Vec<T>`, one `memcpy` per
|
|
/// contiguous run, with no zero-filling of the output first.
|
|
///
|
|
/// For `T` other than `u8`, `elem_size` must equal `size_of::<T>()`. The
|
|
/// selection must be a validated hyperslab, point list or `None` (`All` is the
|
|
/// caller's to handle); `src` must hold exactly the dataset. Anything that
|
|
/// would read outside `src` is an error, never a partial result.
|
|
pub(crate) fn gather<T: NativeElement>(
|
|
src: &[u8],
|
|
dims: &[u64],
|
|
elem_size: usize,
|
|
selection: &Selection,
|
|
) -> Result<Vec<T>, FormatError> {
|
|
let t_size = core::mem::size_of::<T>();
|
|
if elem_size == 0 || (t_size != 1 && t_size != elem_size) {
|
|
return Err(FormatError::DataSizeMismatch {
|
|
expected: t_size,
|
|
actual: elem_size,
|
|
});
|
|
}
|
|
let n_elements = match selection {
|
|
Selection::None => 0,
|
|
Selection::Hyperslab { count, block, .. } => count
|
|
.iter()
|
|
.zip(block)
|
|
.try_fold(1u64, |acc, (&c, &b)| acc.checked_mul(c.checked_mul(b)?))
|
|
.ok_or_else(|| FormatError::Overflow("hyperslab count x block overflows".into()))?,
|
|
Selection::Points(points) => points.len() as u64,
|
|
Selection::All => {
|
|
return Err(FormatError::SelectionOutOfBounds(
|
|
"gather does not take Selection::All".into(),
|
|
));
|
|
}
|
|
};
|
|
let out_bytes = crate::chunked_read::checked_byte_len(n_elements, elem_size)?;
|
|
let out_len = out_bytes / t_size;
|
|
let mut out: Vec<T> = crate::bulk_alloc::vec_for_bulk(out_len);
|
|
let dst = out.as_mut_ptr().cast::<u8>();
|
|
let mut written = 0usize;
|
|
let mut failed = false;
|
|
let mut copy_run = |first: u64, n: u64| {
|
|
if failed {
|
|
return;
|
|
}
|
|
let range = usize::try_from(first)
|
|
.ok()
|
|
.and_then(|f| f.checked_mul(elem_size))
|
|
.zip(
|
|
usize::try_from(n)
|
|
.ok()
|
|
.and_then(|n| n.checked_mul(elem_size)),
|
|
)
|
|
.and_then(|(at, len)| Some((at, len, at.checked_add(len)?)));
|
|
match range {
|
|
Some((at, len, end)) if end <= src.len() && written + len <= out_bytes => {
|
|
// SAFETY: `src[at..end]` is in bounds (checked above), and
|
|
// `dst + written .. + len` lies within `out`'s capacity of
|
|
// `out_bytes` bytes (checked above); `out` is a fresh
|
|
// allocation, so the regions do not overlap.
|
|
unsafe {
|
|
core::ptr::copy_nonoverlapping(src.as_ptr().add(at), dst.add(written), len)
|
|
};
|
|
written += len;
|
|
}
|
|
_ => failed = true,
|
|
}
|
|
};
|
|
let mut bad_point = false;
|
|
match selection {
|
|
Selection::Hyperslab {
|
|
start,
|
|
stride,
|
|
count,
|
|
block,
|
|
} => {
|
|
let rank = dims.len();
|
|
if [start.len(), stride.len(), count.len(), block.len()] != [rank; 4] {
|
|
return Err(FormatError::SelectionOutOfBounds(
|
|
"hyperslab rank does not match dataset rank".into(),
|
|
));
|
|
}
|
|
hyperslab_runs(dims, start, stride, count, block, &mut copy_run);
|
|
}
|
|
Selection::Points(points) => {
|
|
let strides = strides(dims);
|
|
let mut runs = Coalesce {
|
|
start: 0,
|
|
len: 0,
|
|
emit: &mut copy_run,
|
|
};
|
|
for p in points {
|
|
if p.len() != dims.len() || p.iter().zip(dims).any(|(c, n)| c >= n) {
|
|
bad_point = true;
|
|
break;
|
|
}
|
|
let at = p
|
|
.iter()
|
|
.zip(&strides)
|
|
.fold(0u64, |acc, (c, s)| acc.wrapping_add(c.wrapping_mul(*s)));
|
|
runs.push(at, 1);
|
|
}
|
|
runs.flush();
|
|
}
|
|
Selection::None | Selection::All => {}
|
|
}
|
|
if failed || bad_point || written != out_bytes {
|
|
return Err(FormatError::SelectionOutOfBounds(
|
|
"selection addresses elements outside the dataset".into(),
|
|
));
|
|
}
|
|
// SAFETY: all `out_bytes` bytes, i.e. `out_len` values of `T`, were
|
|
// written above, and every bit pattern is a valid `T` (`NativeElement`).
|
|
unsafe { out.set_len(out_len) };
|
|
Ok(out)
|
|
}
|
|
|
|
/// Largest gap between two of a selection's runs that [`gather_storage`]
|
|
/// reads through rather than asking for the runs separately: skipping a
|
|
/// few KiB costs a remote backend far less than another request (and a
|
|
/// local one less than another call and allocation).
|
|
pub(crate) const GATHER_GAP_BYTES: usize = 4 << 10;
|
|
|
|
/// Largest single read [`gather_storage`] makes of a selection's runs: runs
|
|
/// are merged into reads up to this size, and a longer run is split.
|
|
pub(crate) const GATHER_SPAN_BYTES: usize = 8 << 20;
|
|
|
|
/// Call `emit(first_element, element_count)` for each run of a validated
|
|
/// hyperslab or point selection (in output order; see [`hyperslab_runs`]),
|
|
/// or the error for a hyperslab of the wrong rank or a point outside `dims`
|
|
/// (runs before that point have been emitted).
|
|
fn selection_runs(
|
|
dims: &[u64],
|
|
selection: &Selection,
|
|
emit: &mut dyn FnMut(u64, u64),
|
|
) -> Result<(), FormatError> {
|
|
match selection {
|
|
Selection::Hyperslab {
|
|
start,
|
|
stride,
|
|
count,
|
|
block,
|
|
} => {
|
|
let rank = dims.len();
|
|
if [start.len(), stride.len(), count.len(), block.len()] != [rank; 4] {
|
|
return Err(FormatError::SelectionOutOfBounds(
|
|
"hyperslab rank does not match dataset rank".into(),
|
|
));
|
|
}
|
|
hyperslab_runs(dims, start, stride, count, block, emit);
|
|
}
|
|
Selection::Points(points) => {
|
|
let strides = strides(dims);
|
|
let mut coalesce = Coalesce {
|
|
start: 0,
|
|
len: 0,
|
|
emit,
|
|
};
|
|
for p in points {
|
|
if p.len() != dims.len() || p.iter().zip(dims).any(|(c, n)| c >= n) {
|
|
return Err(FormatError::SelectionOutOfBounds(
|
|
"selection addresses elements outside the dataset".into(),
|
|
));
|
|
}
|
|
let at = p
|
|
.iter()
|
|
.zip(&strides)
|
|
.fold(0u64, |acc, (c, s)| acc.wrapping_add(c.wrapping_mul(*s)));
|
|
coalesce.push(at, 1);
|
|
}
|
|
coalesce.flush();
|
|
}
|
|
Selection::None | Selection::All => {}
|
|
}
|
|
Ok(())
|
|
}
|
|
|
|
/// One read of [`gather_storage`]: bytes `[start, end)` of the dataset,
|
|
/// which hold the output's bytes up to `out_end` (from where the previous
|
|
/// span's end left off).
|
|
#[derive(Clone, Copy)]
|
|
struct Span {
|
|
start: usize,
|
|
end: usize,
|
|
out_end: usize,
|
|
}
|
|
|
|
/// [`gather`] of bytes (`T = u8`) from a dataset that is not in memory: the
|
|
/// dataset's `src_len` bytes start at `base` in `file`, which must hold all
|
|
/// of them (the caller checks). Same checks and errors as [`gather`].
|
|
///
|
|
/// The selection's runs are walked twice. The first walk checks them and
|
|
/// plans the reads: runs in increasing order with at most
|
|
/// [`GATHER_GAP_BYTES`] between them are read as one span (the gap is read
|
|
/// and dropped), up to [`GATHER_SPAN_BYTES`] per span. So a strided
|
|
/// selection is a few large reads, not one per element, and nothing is
|
|
/// allocated per run. The spans are fetched batch by batch (one
|
|
/// [`Storage::read_ranges`] call per [`crate::storage::RAW_BATCH_BYTES`])
|
|
/// while the second walk copies each run out of its span.
|
|
pub(crate) fn gather_storage<S: Storage + ?Sized>(
|
|
file: &S,
|
|
base: u64,
|
|
src_len: usize,
|
|
dims: &[u64],
|
|
elem_size: usize,
|
|
selection: &Selection,
|
|
) -> Result<Vec<u8>, FormatError> {
|
|
if elem_size == 0 {
|
|
return Err(FormatError::DataSizeMismatch {
|
|
expected: 1,
|
|
actual: elem_size,
|
|
});
|
|
}
|
|
let n_elements = match selection {
|
|
Selection::None => 0,
|
|
Selection::Hyperslab { count, block, .. } => count
|
|
.iter()
|
|
.zip(block)
|
|
.try_fold(1u64, |acc, (&c, &b)| acc.checked_mul(c.checked_mul(b)?))
|
|
.ok_or_else(|| FormatError::Overflow("hyperslab count x block overflows".into()))?,
|
|
Selection::Points(points) => points.len() as u64,
|
|
Selection::All => {
|
|
return Err(FormatError::SelectionOutOfBounds(
|
|
"gather does not take Selection::All".into(),
|
|
));
|
|
}
|
|
};
|
|
let out_bytes = crate::chunked_read::checked_byte_len(n_elements, elem_size)?;
|
|
let outside = || {
|
|
FormatError::SelectionOutOfBounds("selection addresses elements outside the dataset".into())
|
|
};
|
|
|
|
// First walk: check every run and plan the spans.
|
|
let mut spans: Vec<Span> = Vec::new();
|
|
let mut total = 0usize;
|
|
let mut failed = false;
|
|
selection_runs(dims, selection, &mut |first: u64, n: u64| {
|
|
if failed {
|
|
return;
|
|
}
|
|
let range = usize::try_from(first)
|
|
.ok()
|
|
.and_then(|f| f.checked_mul(elem_size))
|
|
.zip(
|
|
usize::try_from(n)
|
|
.ok()
|
|
.and_then(|n| n.checked_mul(elem_size)),
|
|
)
|
|
.and_then(|(at, len)| Some((at, len, at.checked_add(len)?)));
|
|
let Some((mut at, mut len)) = range
|
|
.filter(|&(_, len, end)| end <= src_len && len <= out_bytes - total)
|
|
.map(|(at, len, _)| (at, len))
|
|
else {
|
|
failed = true;
|
|
return;
|
|
};
|
|
while len > 0 {
|
|
let room = match spans.last_mut() {
|
|
Some(s)
|
|
if at >= s.end
|
|
&& at - s.end <= GATHER_GAP_BYTES
|
|
&& at - s.start < GATHER_SPAN_BYTES =>
|
|
{
|
|
let take = len.min(GATHER_SPAN_BYTES - (at - s.start));
|
|
s.end = at + take;
|
|
s.out_end += take;
|
|
take
|
|
}
|
|
_ => {
|
|
let take = len.min(GATHER_SPAN_BYTES);
|
|
spans.push(Span {
|
|
start: at,
|
|
end: at + take,
|
|
out_end: total + take,
|
|
});
|
|
take
|
|
}
|
|
};
|
|
total += room;
|
|
at += room;
|
|
len -= room;
|
|
}
|
|
})?;
|
|
if failed || total != out_bytes {
|
|
return Err(outside());
|
|
}
|
|
|
|
// The spans' reads, and the batches they are fetched in.
|
|
let reqs: Vec<ExtentReq> = spans
|
|
.iter()
|
|
.map(|s| ExtentReq {
|
|
addr: base + s.start as u64,
|
|
len: s.end - s.start,
|
|
fetch: Some(s.end - s.start),
|
|
})
|
|
.collect();
|
|
let batches = raw_batches(reqs.len(), false, |i| reqs[i].len);
|
|
|
|
// Second walk: copy each run out of its span, fetching each batch of
|
|
// spans when the walk reaches it (and dropping the previous one).
|
|
let mut out = crate::bulk_alloc::vec_for_bulk(out_bytes);
|
|
let mut span = 0usize;
|
|
let mut batch = 0usize;
|
|
let mut fetched: Option<ExtentBytes<'_>> = None;
|
|
let mut error: Option<FormatError> = None;
|
|
selection_runs(dims, selection, &mut |first: u64, n: u64| {
|
|
if error.is_some() {
|
|
return;
|
|
}
|
|
// Checked by the first walk (these cannot saturate or wrap).
|
|
let mut at = crate::addr::saturating_usize(first).wrapping_mul(elem_size);
|
|
let mut len = crate::addr::saturating_usize(n).wrapping_mul(elem_size);
|
|
while len > 0 {
|
|
while spans.get(span).is_some_and(|s| s.out_end <= out.len()) {
|
|
span += 1;
|
|
}
|
|
if fetched.is_none() || span >= batches[batch].end {
|
|
fetched = None;
|
|
while batches.get(batch).is_some_and(|b| span >= b.end) {
|
|
batch += 1;
|
|
}
|
|
let (Some(b), Some(_)) = (batches.get(batch).cloned(), spans.get(span)) else {
|
|
// The second walk emitted more than the first.
|
|
error = Some(outside());
|
|
return;
|
|
};
|
|
match ExtentBytes::fetch(file, &reqs[b.clone()], b.start) {
|
|
Ok(f) => fetched = Some(f),
|
|
Err(e) => {
|
|
error = Some(e);
|
|
return;
|
|
}
|
|
}
|
|
}
|
|
let s = spans[span];
|
|
let take = len.min(s.out_end - out.len());
|
|
let bytes = match fetched
|
|
.as_ref()
|
|
.map(|f| f.get(span, &reqs[span]))
|
|
.unwrap_or_else(|| Err(outside()))
|
|
{
|
|
Ok(b) => b,
|
|
Err(e) => {
|
|
error = Some(e);
|
|
return;
|
|
}
|
|
};
|
|
match at
|
|
.checked_sub(s.start)
|
|
.and_then(|o| bytes.get(o..o.checked_add(take)?))
|
|
{
|
|
Some(b) => out.extend_from_slice(b),
|
|
None => {
|
|
error = Some(outside());
|
|
return;
|
|
}
|
|
}
|
|
at += take;
|
|
len -= take;
|
|
}
|
|
})?;
|
|
if let Some(e) = error {
|
|
return Err(e);
|
|
}
|
|
if out.len() != out_bytes {
|
|
return Err(outside());
|
|
}
|
|
Ok(out)
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
fn runs(dims: &[u64], sel: [&[u64]; 4]) -> Vec<(u64, u64)> {
|
|
let mut v = Vec::new();
|
|
hyperslab_runs(dims, sel[0], sel[1], sel[2], sel[3], |s, n| v.push((s, n)));
|
|
v
|
|
}
|
|
|
|
#[test]
|
|
fn runs_merge_blocks_and_whole_rows() {
|
|
// A box: one run per row.
|
|
assert_eq!(
|
|
runs(&[4, 10], [&[1, 2], &[1, 1], &[2, 3], &[1, 1]]),
|
|
vec![(12, 3), (22, 3)]
|
|
);
|
|
// Whole rows: one run.
|
|
assert_eq!(
|
|
runs(&[4, 10], [&[1, 0], &[1, 1], &[3, 10], &[1, 1]]),
|
|
vec![(10, 30)]
|
|
);
|
|
// stride == block: blocks merge.
|
|
assert_eq!(
|
|
runs(&[1, 10], [&[0, 1], &[1, 2], &[1, 4], &[1, 2]]),
|
|
vec![(1, 8)]
|
|
);
|
|
// Strided with blocks along both dimensions.
|
|
assert_eq!(
|
|
runs(&[6, 10], [&[0, 1], &[3, 4], &[2, 2], &[2, 2]]),
|
|
vec![
|
|
(1, 2),
|
|
(5, 2),
|
|
(11, 2),
|
|
(15, 2),
|
|
(31, 2),
|
|
(35, 2),
|
|
(41, 2),
|
|
(45, 2)
|
|
]
|
|
);
|
|
// Empty.
|
|
assert!(runs(&[4, 10], [&[0, 0], &[1, 1], &[0, 3], &[1, 1]]).is_empty());
|
|
// Scalar.
|
|
assert_eq!(runs(&[], [&[], &[], &[], &[]]), vec![(0, 1)]);
|
|
}
|
|
|
|
#[test]
|
|
fn gather_matches_element_order_and_rejects_out_of_range() {
|
|
let dims = [3u64, 4];
|
|
let src: Vec<u8> = (0..12u16).flat_map(|v| v.to_le_bytes()).collect();
|
|
let sel = Selection::Hyperslab {
|
|
start: vec![0, 1],
|
|
stride: vec![2, 2],
|
|
count: vec![2, 2],
|
|
block: vec![1, 1],
|
|
};
|
|
let got: Vec<u8> = gather(&src, &dims, 2, &sel).unwrap();
|
|
let want: Vec<u8> = [1u16, 3, 9, 11]
|
|
.iter()
|
|
.flat_map(|v| v.to_le_bytes())
|
|
.collect();
|
|
assert_eq!(got, want);
|
|
let pts = Selection::Points(vec![vec![2, 3], vec![0, 0], vec![0, 1]]);
|
|
let got: Vec<u8> = gather(&src, &dims, 2, &pts).unwrap();
|
|
let want: Vec<u8> = [11u16, 0, 1].iter().flat_map(|v| v.to_le_bytes()).collect();
|
|
assert_eq!(got, want);
|
|
// Past the extent, or a source shorter than the dataset: an error.
|
|
let bad = Selection::Points(vec![vec![3, 0]]);
|
|
assert!(gather::<u8>(&src, &dims, 2, &bad).is_err());
|
|
let past = Selection::Hyperslab {
|
|
start: vec![2, 0],
|
|
stride: vec![1, 1],
|
|
count: vec![2, 4],
|
|
block: vec![1, 1],
|
|
};
|
|
assert!(gather::<u8>(&src, &dims, 2, &past).is_err());
|
|
assert!(gather::<u8>(&src[..20], &dims, 2, &pts).is_err());
|
|
}
|
|
}
|