use onnx_runtime_ep_api::{EpError, Result};
pub fn next_index(shape: &[usize], index: &mut [usize]) -> bool {
for axis in (0..shape.len()).rev() {
index[axis] += 1;
if index[axis] < shape[axis] {
return true;
}
index[axis] = 0;
}
false
}
pub fn elem_offset(strides: &[i64], index: &[usize]) -> isize {
let mut offset = 0i64;
for (stride, &i) in strides.iter().zip(index) {
offset += stride * i as i64;
}
offset as isize
}
pub fn numel(shape: &[usize]) -> usize {
shape.iter().product()
}
pub fn addressed_elem_range(shape: &[usize], strides: &[i64]) -> (i64, i64) {
let mut min = 0i64;
let mut max = 0i64;
for (&dim, &stride) in shape.iter().zip(strides) {
if dim == 0 {
continue;
}
let extent = (dim as i64 - 1) * stride;
if extent < 0 {
min += extent;
} else {
max += extent;
}
}
(min, max)
}
pub fn view_in_bounds(
shape: &[usize],
strides: &[i64],
byte_offset: usize,
esize: usize,
buffer_len: usize,
) -> Result<()> {
if shape.len() != strides.len() {
return Err(EpError::InvalidTensorView {
reason: format!(
"rank mismatch: shape {} dims, strides {}",
shape.len(),
strides.len()
),
});
}
if numel(shape) == 0 {
return Ok(());
}
let (min_elem, max_elem) = addressed_elem_range(shape, strides);
let esize = esize as i128;
let origin = byte_offset as i128;
let lo = (min_elem as i128)
.checked_mul(esize)
.and_then(|m| origin.checked_add(m));
let hi = (max_elem as i128)
.checked_mul(esize)
.and_then(|m| origin.checked_add(m))
.and_then(|h| h.checked_add(esize)); let (lo, hi) = match (lo, hi) {
(Some(lo), Some(hi)) => (lo, hi),
_ => {
return Err(EpError::InvalidTensorView {
reason: format!(
"view address computation overflowed (shape {shape:?}, strides {strides:?}, \
byte_offset {byte_offset}, esize {esize})"
),
});
}
};
if lo < 0 || hi > buffer_len as i128 {
return Err(EpError::InvalidTensorView {
reason: format!(
"view addresses bytes [{lo}, {hi}) outside backing allocation [0, {buffer_len})"
),
});
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn next_index_walks_row_major() {
let shape = [2usize, 3];
let mut idx = [0usize, 0];
let mut seen = vec![idx];
while next_index(&shape, &mut idx) {
seen.push(idx);
}
assert_eq!(seen.len(), 6);
assert_eq!(seen[0], [0, 0]);
assert_eq!(seen[1], [0, 1]);
assert_eq!(seen[3], [1, 0]);
assert_eq!(seen[5], [1, 2]);
}
#[test]
fn scalar_has_single_element() {
let shape: [usize; 0] = [];
let mut idx: [usize; 0] = [];
assert!(!next_index(&shape, &mut idx));
}
#[test]
fn offset_uses_strides() {
assert_eq!(elem_offset(&[1, 3], &[0, 0]), 0);
assert_eq!(elem_offset(&[1, 3], &[2, 1]), 2 + 3);
}
#[test]
fn addressed_range_handles_negative_strides() {
assert_eq!(addressed_elem_range(&[4], &[-1]), (-3, 0));
assert_eq!(addressed_elem_range(&[2, 3], &[3, 1]), (0, 5));
}
#[test]
fn bounds_accept_contiguous_and_reject_overrun() {
assert!(view_in_bounds(&[2, 3], &[3, 1], 0, 4, 24).is_ok());
assert!(view_in_bounds(&[2, 3], &[3, 1], 0, 4, 23).is_err());
}
#[test]
fn bounds_reject_negative_stride_underrun() {
assert!(view_in_bounds(&[4], &[-1], 0, 4, 16).is_err());
assert!(view_in_bounds(&[4], &[-1], 12, 4, 16).is_ok());
}
#[test]
fn empty_tensor_is_in_bounds() {
assert!(view_in_bounds(&[0, 5], &[5, 1], 0, 4, 0).is_ok());
}
#[test]
fn bounds_reject_overflowing_address_math() {
let shape = [2usize];
let strides = [i64::MAX];
let err = view_in_bounds(&shape, &strides, 0, 8, 1024);
assert!(
err.is_err(),
"an address computation that overflows must be rejected, never wrap-passed"
);
}
}