use crate::{MattenError, Tensor};
#[test]
fn builder_all_all() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let s = t.slice().all().all().build().unwrap();
assert_eq!(s, t);
}
#[test]
fn builder_index_first_row() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let row = t.slice().index(0).all().build().unwrap();
assert_eq!(row.shape(), &[3]);
assert_eq!(row.as_slice(), &[1.0, 2.0, 3.0]);
}
#[test]
fn builder_index_second_row() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let row = t.slice().index(1).all().build().unwrap();
assert_eq!(row.shape(), &[3]);
assert_eq!(row.as_slice(), &[4.0, 5.0, 6.0]);
}
#[test]
fn builder_range_rows() {
let t = Tensor::new((1..=12).map(|x| x as f64).collect(), &[3, 4]);
let s = t.slice().range(0..2).all().build().unwrap();
assert_eq!(s.shape(), &[2, 4]);
assert_eq!(s.as_slice(), &[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]);
}
#[test]
fn builder_range_cols() {
let t = Tensor::new((1..=6).map(|x| x as f64).collect(), &[2, 3]);
let s = t.slice().all().range(1..3).build().unwrap();
assert_eq!(s.shape(), &[2, 2]);
assert_eq!(s.as_slice(), &[2.0, 3.0, 5.0, 6.0]);
}
#[test]
fn builder_range_from() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[4]);
let s = t.slice().range(2..).build().unwrap();
assert_eq!(s.as_slice(), &[3.0, 4.0]);
}
#[test]
fn builder_range_to() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[4]);
let s = t.slice().range(..2).build().unwrap();
assert_eq!(s.as_slice(), &[1.0, 2.0]);
}
#[test]
fn builder_range_full() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[4]);
let s = t.slice().range(..).build().unwrap();
assert_eq!(s, t);
}
#[test]
fn builder_inclusive_range() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[4]);
let s = t.slice().range(1..=2).build().unwrap();
assert_eq!(s.as_slice(), &[2.0, 3.0]);
}
#[test]
fn builder_index_all_axes_gives_scalar() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[2, 2]);
let s = t.slice().index(1).index(0).build().unwrap();
assert!(s.is_scalar());
assert_eq!(s.as_slice(), &[3.0]);
}
#[test]
fn builder_rank_mismatch_is_err() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[2, 2]);
let err = t.slice().all().build().unwrap_err();
assert!(matches!(err, MattenError::Slice { .. }));
}
#[test]
fn builder_out_of_bounds_index_is_err() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[2, 2]);
let err = t.slice().index(5).all().build().unwrap_err();
assert!(matches!(err, MattenError::Slice { .. }));
}
#[test]
fn builder_result_is_independent() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let s = t.slice().index(0).all().build().unwrap();
assert_eq!(t.len(), 6);
assert_eq!(s.len(), 3);
}
#[test]
fn slice_str_all() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[2, 2]);
let s = t.slice_str(":, :").unwrap();
assert_eq!(s, t);
}
#[test]
fn slice_str_first_row() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let s = t.slice_str("0, :").unwrap();
assert_eq!(s.shape(), &[3]);
assert_eq!(s.as_slice(), &[1.0, 2.0, 3.0]);
}
#[test]
fn slice_str_range() {
let t = Tensor::new((1..=12).map(|x| x as f64).collect(), &[3, 4]);
let s = t.slice_str("0:2, :").unwrap();
assert_eq!(s.shape(), &[2, 4]);
}
#[test]
fn slice_str_range_from() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[4]);
let s = t.slice_str("2:").unwrap();
assert_eq!(s.as_slice(), &[3.0, 4.0]);
}
#[test]
fn slice_str_range_to() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[4]);
let s = t.slice_str(":2").unwrap();
assert_eq!(s.as_slice(), &[1.0, 2.0]);
}
#[test]
fn slice_str_step() {
let t = Tensor::new((0..=9).map(|x| x as f64).collect(), &[10]);
let s = t.slice_str("0:10:2").unwrap();
assert_eq!(s.as_slice(), &[0.0, 2.0, 4.0, 6.0, 8.0]);
}
#[test]
fn slice_str_whitespace_ignored() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let a = t.slice_str("0,:").unwrap();
let b = t.slice_str(" 0 , : ").unwrap();
assert_eq!(a, b);
}
#[test]
fn slice_str_matches_builder() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[2, 3]);
let from_str = t.slice_str("0:2, :").unwrap();
let from_builder = t.slice().range(0..2).all().build().unwrap();
assert_eq!(from_str, from_builder);
}
#[test]
fn slice_str_malformed_is_err() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[2, 2]);
for bad in &["0::", "a:b", ":::", "", "x"] {
assert!(
t.slice_str(bad).is_err(),
"expected Err for {:?} but got Ok",
bad
);
}
}
#[test]
fn slice_str_too_many_dims_is_err() {
let t = Tensor::new(vec![1.0, 2.0, 3.0, 4.0], &[2, 2]);
let err = t.slice_str("0, 0, 0").unwrap_err();
assert!(matches!(err, MattenError::Slice { .. }));
}
#[test]
fn slice_str_oversized_is_err() {
let t = Tensor::new(vec![1.0, 2.0], &[2]);
let long = "0:1, ".repeat(200);
let err = t.slice_str(&long).unwrap_err();
assert!(matches!(err, MattenError::Slice { .. }));
assert!(err.to_string().contains("maximum length"));
}
#[test]
fn negative_index_and_range_forms_on_a_3_element_vector() {
let t = Tensor::new(vec![1.0, 2.0, 3.0], &[3]);
assert_eq!(t.slice_str("-1").unwrap().as_slice(), &[3.0]);
assert_eq!(t.slice_str("-2").unwrap().as_slice(), &[2.0]);
assert_eq!(t.slice_str("-3").unwrap().as_slice(), &[1.0]); assert_eq!(t.slice_str("0:-1").unwrap().as_slice(), &[1.0, 2.0]);
assert_eq!(t.slice_str(":-1").unwrap().as_slice(), &[1.0, 2.0]);
assert_eq!(t.slice_str("-2:").unwrap().as_slice(), &[2.0, 3.0]);
assert_eq!(t.slice_str("-3:-1").unwrap().as_slice(), &[1.0, 2.0]);
}
#[test]
fn negative_index_on_every_axis_of_an_unequal_rank2_tensor() {
let m = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[3, 2]);
let last_row = m.slice_str("-1,:").unwrap();
assert_eq!(last_row.shape(), &[2]);
assert_eq!(last_row.as_slice(), &[5.0, 6.0]);
let last_col = m.slice_str(":,-1").unwrap();
assert_eq!(last_col.shape(), &[3]);
assert_eq!(last_col.as_slice(), &[2.0, 4.0, 6.0]);
let corner = m.slice_str("-1,-1").unwrap();
assert!(corner.is_scalar());
assert_eq!(corner.as_slice(), &[6.0]);
}
#[test]
fn negative_index_mixed_signs_in_one_spec() {
let m = Tensor::new(vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &[3, 2]);
let s = m.slice_str("0:-1,-1").unwrap();
assert_eq!(s.shape(), &[2]);
assert_eq!(s.as_slice(), &[2.0, 4.0]);
}
#[test]
fn negative_index_n_equals_dim_plus_1_errors() {
let t = Tensor::new(vec![1.0, 2.0, 3.0], &[3]);
assert!(t.slice_str("-3").is_ok()); assert!(t.slice_str("-4").is_err()); }
#[test]
fn negative_out_of_range_errors_rather_than_clamps() {
let t = Tensor::new(vec![1.0, 2.0, 3.0], &[3]);
let err_index = t.slice_str("-10").unwrap_err();
assert!(matches!(err_index, MattenError::Slice { .. }));
let err_range = t.slice_str("-10:").unwrap_err();
assert!(matches!(err_range, MattenError::Slice { .. }));
assert_ne!(
t.slice_str("-10:").ok().map(|r| r.as_slice().to_vec()),
Some(vec![1.0, 2.0, 3.0])
);
}
#[test]
fn negative_out_of_range_error_message_shows_written_and_resolved_forms() {
let t = Tensor::new(vec![1.0, 2.0, 3.0], &[3]);
let err = t.slice_str("-10").unwrap_err();
let message = err.to_string();
assert!(
message.contains("-10") && message.contains("-7"),
"message should show both the written form (-10) and the resolved index (-7): {message}"
);
}
#[test]
fn inverted_range_message_names_written_forms_when_negative() {
let t = Tensor::new(vec![1.0, 2.0, 3.0], &[3]);
let err = t.slice_str("-1:-3").unwrap_err();
let message = err.to_string();
assert!(
message.contains("-1")
&& message.contains("-3")
&& message.contains('2')
&& message.contains('0'),
"message should name both written forms (-1, -3) and both resolutions (2, 0): {message}"
);
let mixed_err = t.slice_str("2:-3").unwrap_err();
let mixed_message = mixed_err.to_string();
assert!(
mixed_message.contains("-3") && mixed_message.contains('0'),
"message should name the written negative end (-3) and its resolution (0): {mixed_message}"
);
}
#[test]
fn inverted_range_message_unchanged_when_both_bounds_non_negative() {
let t = Tensor::new(vec![1.0, 2.0, 3.0], &[3]);
let err = t.slice_str("2:1").unwrap_err();
assert_eq!(
err.to_string(),
"matten slice error: range start 2 > end 1 for axis 0 in slice_str"
);
}
#[test]
fn negative_step_reversal_is_still_a_parse_error() {
let t = Tensor::new(vec![1.0, 2.0, 3.0], &[3]);
assert!(t.slice_str("::-1").is_err());
}
#[test]
fn negative_zero_behaves_as_zero() {
let t = Tensor::new(vec![1.0, 2.0, 3.0], &[3]);
assert_eq!(
t.slice_str("-0").unwrap().as_slice(),
t.slice_str("0").unwrap().as_slice()
);
}
#[test]
fn preexisting_specs_still_parse_to_identical_results() {
let t = Tensor::new((0..=9).map(|x| x as f64).collect(), &[10]);
assert_eq!(t.slice_str(":").unwrap(), t.slice_str(":").unwrap());
assert_eq!(t.slice_str("0").unwrap().as_slice(), &[0.0]);
assert_eq!(t.slice_str("0:2").unwrap().as_slice(), &[0.0, 1.0]);
assert_eq!(
t.slice_str("2:").unwrap().as_slice(),
&[2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0]
);
assert_eq!(t.slice_str(":2").unwrap().as_slice(), &[0.0, 1.0]);
assert_eq!(
t.slice_str("0:10:2").unwrap().as_slice(),
&[0.0, 2.0, 4.0, 6.0, 8.0]
);
}