mod utils;
use boreholeio::strided_array_file::{ArrayProxy, StridedArrayFile};
use std::{path::PathBuf, str::FromStr};
use utils::{assert_eq_views, generate_array};
#[test]
fn self_generated_read_write() {
let arr_i8 = generate_array(&vec![10]);
let arr_u8 = generate_array(&vec![1, 1, 1]);
let arr_u8_2 = generate_array(&vec![2, 3, 1, 4]).t().to_owned();
let arr_i16 = generate_array(&vec![2, 1, 2]);
let arr_u16 = generate_array(&vec![3, 3]);
let arr_u16_2 = generate_array(&vec![2, 2, 2, 2, 2, 2]);
let arr_i32 = generate_array(&vec![5]);
let arr_u32 = generate_array(&vec![1, 1, 1, 1, 3, 4, 1, 2, 1, 1, 1]);
let arr_i64 = generate_array(&vec![4, 4]);
let arr_u64 = generate_array(&vec![3, 3]);
let arr_f32 = generate_array(&vec![2, 4]);
let arr_f64 = generate_array(&vec![2, 2]);
let gt = vec![
ArrayProxy::I8(arr_i8.view()),
ArrayProxy::I8(arr_i8.view()), ArrayProxy::U8(arr_u8.view()),
ArrayProxy::I32(arr_i32.view()), ArrayProxy::U8(arr_u8_2.view()), ArrayProxy::U16(arr_u16_2.view()), ArrayProxy::U16(arr_u16.view()),
ArrayProxy::F64(arr_f64.view()),
ArrayProxy::U32(arr_u32.view()),
ArrayProxy::I16(arr_i16.view()),
ArrayProxy::I64(arr_i64.view()),
ArrayProxy::U64(arr_u64.view()),
ArrayProxy::F32(arr_f32.view()),
];
let file_path = path("1.star");
StridedArrayFile::write(&file_path, >).unwrap();
let ours = StridedArrayFile::new(file_path);
let mut i = 0;
i += assert_eq_views::<i8>(>[i], &ours.get(i));
i += assert_eq_views::<i8>(>[i], &ours.get(i));
i += assert_eq_views::<u8>(>[i], &ours.get(i));
i += assert_eq_views::<i32>(>[i], &ours.get(i));
i += assert_eq_views::<u8>(>[i], &ours.get(i));
i += assert_eq_views::<u16>(>[i], &ours.get(i));
i += assert_eq_views::<u16>(>[i], &ours.get(i));
i += assert_eq_views::<f64>(>[i], &ours.get(i));
i += assert_eq_views::<u32>(>[i], &ours.get(i));
i += assert_eq_views::<i16>(>[i], &ours.get(i));
i += assert_eq_views::<i64>(>[i], &ours.get(i));
i += assert_eq_views::<u64>(>[i], &ours.get(i));
i += assert_eq_views::<f32>(>[i], &ours.get(i));
assert_eq!(i, gt.len());
}
#[test]
fn python_generated_read() {
let file_path = path("tests/strided_array_file/fixture/simple.star");
let arrays = StridedArrayFile::new(file_path);
assert_eq!(arrays.count().unwrap(), 4);
assert_eq!(arrays.get(0).shape(), [3, 4, 2]);
assert_eq!(arrays.get(0).element_type(), 9);
let iden3f64 = ndarray::Array2::from_diag(&ndarray::arr1(&[1u8, 1, 1])).into_dyn();
let view = ArrayProxy::U8(iden3f64.view());
assert_eq_views::<u8>(&arrays.get(1), &view);
assert_eq!(arrays.get(2).shape(), [1, 1, 1, 1, 2, 3, 1, 1, 4, 1, 1]);
assert_eq!(arrays.get(2).element_type(), 4);
let zeros10f32 = ndarray::Array1::from_elem([10], 0.0f32).into_dyn();
let view = ArrayProxy::F32(zeros10f32.view());
assert_eq_views::<f32>(&arrays.get(3), &view);
}
#[test]
fn verify_strides() {
let data = generate_array(&vec![12]);
let proxy = ArrayProxy::F64(data.view());
let step = 2;
let slice_info = ndarray::SliceInfoElem::Slice {
start: 0,
end: None,
step: step as isize,
};
let slice = proxy.slice(vec![slice_info]).unwrap();
let path = path("slice.star");
StridedArrayFile::write(&path, &vec![slice.clone()]).unwrap();
let file = StridedArrayFile::new(path);
let view = file.get_as::<f64>(0);
for i in 0..view.len() {
assert_eq!(view[i], data[i * step]);
}
}
fn path(path: &str) -> PathBuf {
PathBuf::from_str(path).unwrap()
}