diff --git a/diskann-utils/src/views/rowmajor.rs b/diskann-utils/src/views/rowmajor.rs index ffd569394d..c32f3748a2 100644 --- a/diskann-utils/src/views/rowmajor.rs +++ b/diskann-utils/src/views/rowmajor.rs @@ -519,6 +519,27 @@ pub unsafe trait MatrixMut: Matrix { } } + /// Return a mutable view over the rows in `rows`, or `None` when the range is invalid. + fn subview_mut(&mut self, rows: std::ops::Range) -> Option> { + if rows.start > rows.end || rows.end > self.nrows() { + return None; + } + + let ncols = self.ncols(); + // SAFETY: `rows.start <= self.nrows()`, so the validated parent layout makes the + // offset representable and places it within or one past the matrix span. Deriving + // from `as_nonnull_mut` preserves the mutable provenance required by `MatrixMut`. + let ptr = unsafe { self.as_nonnull_mut().add(rows.start * ncols) }; + // SAFETY: The selected rows are a subset of the validated parent layout, with the + // same column count. + let layout = unsafe { Layout::new_unchecked(rows.end - rows.start, ncols) }; + Some(Mut { + ptr, + layout, + _lifetime: PhantomData, + }) + } + //-------// // Rayon // //-------// @@ -2207,14 +2228,25 @@ mod tests { assert_eq!(m.get_element_mut(0, 3), None); } + /// Return `m.subview(rows)` after checking that `m.subview_mut(rows)` selects the same + /// span of `m`. + fn check_subviews(m: &mut Owned, rows: std::ops::Range) -> Option> { + let expected = m + .subview_mut(rows.clone()) + .map(|v| (v.as_ptr(), v.layout())); + let subview = m.subview(rows); + assert_eq!(subview.map(|v| (v.as_ptr(), v.layout())), expected); + subview + } + #[test] - fn test_subview() { + fn test_subviews() { let data = make_test_matrix(); - let m = Owned::try_from_data(data.into(), 4, 3).unwrap(); + let mut m = Owned::try_from_data(data.into(), 4, 3).unwrap(); // Create a subview of the first two rows { - let subview = m.subview(0..4).unwrap(); + let subview = check_subviews(&mut m, 0..4).unwrap(); assert_eq!(subview.nrows(), 4); assert_eq!(subview.ncols(), 3); @@ -2227,7 +2259,7 @@ mod tests { // Sub view over a subset that touches the end. { - let subview = m.subview(1..4).unwrap(); + let subview = check_subviews(&mut m, 1..4).unwrap(); assert_eq!(subview.nrows(), 3); assert_eq!(subview.ncols(), 3); @@ -2239,27 +2271,27 @@ mod tests { // Empty subview in the middle. { - let subview = m.subview(2..2).unwrap(); + let subview = check_subviews(&mut m, 2..2).unwrap(); assert_eq!(subview.nrows(), 0); assert_eq!(subview.ncols(), 3); } // Empty subviews at both boundaries. { - let subview = m.subview(0..0).unwrap(); + let subview = check_subviews(&mut m, 0..0).unwrap(); assert_eq!(subview.nrows(), 0); assert_eq!(subview.ncols(), 3); - let subview = m.subview(4..4).unwrap(); + let subview = check_subviews(&mut m, 4..4).unwrap(); assert_eq!(subview.nrows(), 0); assert_eq!(subview.ncols(), 3); } // Empty out-of-bounds subview - assert!(m.subview(5..5).is_none()); + assert!(check_subviews(&mut m, 5..5).is_none()); // End is out of bounds. - assert!(m.subview(2..10).is_none()); + assert!(check_subviews(&mut m, 2..10).is_none()); // Reversed bounds. #[expect( @@ -2267,10 +2299,21 @@ mod tests { reason = "we want to make sure it doesn't work" )] let empty = 3..2; - assert!(m.subview(empty).is_none()); + assert!(check_subviews(&mut m, empty).is_none()); // Extreme out-of-bounds values do not overflow internal calculations. - assert!(m.subview(usize::MAX - 1..usize::MAX).is_none()); + assert!(check_subviews(&mut m, usize::MAX - 1..usize::MAX).is_none()); + + // Writes through (nested) mutable subviews land only in the selected parent rows. + { + let mut outer = m.subview_mut(1..4).unwrap(); + outer.row_mut(0).copy_from_slice(&[10, 11, 12]); + + let mut inner = outer.subview_mut(1..2).unwrap(); + assert_eq!(inner.row(0), &[2, 3, 4]); + inner.as_mut_slice().fill(7); + } + assert_eq!(m.as_slice(), &[0, 1, 2, 10, 11, 12, 7, 7, 7, 3, 4, 5]); } #[expect( @@ -2279,29 +2322,29 @@ mod tests { )] #[test] fn test_subview_zero_cols() { - let m = Owned::from_element(10, 0, 0u32); + let mut m = Owned::from_element(10, 0, 0u32); // A fully disjoint range is rejected. - assert!(m.subview(100..200).is_none()); + assert!(check_subviews(&mut m, 100..200).is_none()); // A range extending one row beyond the matrix is rejected. - assert!(m.subview(10..11).is_none()); + assert!(check_subviews(&mut m, 10..11).is_none()); // Reversed bounds are rejected. - assert!(m.subview(5..4).is_none()); + assert!(check_subviews(&mut m, 5..4).is_none()); // An in-bounds suffix preserves its logical shape. - let v = m.subview(5..10).unwrap(); + let v = check_subviews(&mut m, 5..10).unwrap(); assert_eq!(v.nrows(), 5); assert_eq!(v.ncols(), 0); // An empty range at the end is valid. - let v = m.subview(10..10).unwrap(); + let v = check_subviews(&mut m, 10..10).unwrap(); assert_eq!(v.nrows(), 0); assert_eq!(v.ncols(), 0); // The full range preserves all logical rows despite having no elements. - let v = m.subview(0..10).unwrap(); + let v = check_subviews(&mut m, 0..10).unwrap(); assert_eq!(v.nrows(), 10); assert_eq!(v.ncols(), 0); }