1use std::fs::{self, File};
4use std::io::{self, Write};
5use std::path::{Path, PathBuf};
6use std::time::{SystemTime, UNIX_EPOCH};
7
8use ndarray::{Array1, Array2, ArrayView1, Axis};
9
10use crate::error::Result;
11
12pub fn atomic_write_file<F>(path: &Path, write_fn: F) -> Result<()>
17where
18 F: FnOnce(&mut File) -> Result<()>,
19{
20 let parent = path.parent().unwrap_or_else(|| Path::new("."));
21 fs::create_dir_all(parent)?;
22
23 let mut temp_path = atomic_temp_path(path);
24 let mut attempts = 0u32;
25 let mut file = loop {
26 match File::options()
27 .write(true)
28 .create_new(true)
29 .open(&temp_path)
30 {
31 Ok(file) => break file,
32 Err(err) if err.kind() == io::ErrorKind::AlreadyExists && attempts < 8 => {
33 attempts += 1;
34 temp_path = atomic_temp_path_with_attempt(path, attempts);
35 }
36 Err(err) => return Err(err.into()),
37 }
38 };
39
40 let write_result = write_fn(&mut file).and_then(|_| {
41 file.flush()?;
42 file.sync_all()?;
43 Ok(())
44 });
45
46 if let Err(err) = write_result {
47 drop(file);
48 let _ = fs::remove_file(&temp_path);
49 return Err(err);
50 }
51
52 drop(file);
53 fs::rename(&temp_path, path)?;
54
55 if let Ok(parent_dir) = File::open(parent) {
56 let _ = parent_dir.sync_all();
57 }
58
59 Ok(())
60}
61
62fn atomic_temp_path(path: &Path) -> PathBuf {
63 atomic_temp_path_with_attempt(path, 0)
64}
65
66fn atomic_temp_path_with_attempt(path: &Path, attempt: u32) -> PathBuf {
67 let file_name = path
68 .file_name()
69 .and_then(|name| name.to_str())
70 .unwrap_or("atomic-write");
71 let nanos = SystemTime::now()
72 .duration_since(UNIX_EPOCH)
73 .map(|duration| duration.as_nanos())
74 .unwrap_or_default();
75 let tmp_name = format!(
76 ".{}.{}.{}.tmp",
77 file_name,
78 std::process::id(),
79 nanos + attempt as u128
80 );
81 path.with_file_name(tmp_name)
82}
83
84pub fn quantile(arr: &Array1<f32>, q: f64) -> f32 {
95 if arr.is_empty() {
96 return 0.0;
97 }
98
99 let mut sorted: Vec<f32> = arr.iter().copied().collect();
100 sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
101
102 let n = sorted.len();
103 let idx_float = q * (n - 1) as f64;
104 let lower_idx = idx_float.floor() as usize;
105 let upper_idx = idx_float.ceil() as usize;
106
107 if lower_idx == upper_idx {
108 sorted[lower_idx]
109 } else {
110 let weight = (idx_float - lower_idx as f64) as f32;
111 sorted[lower_idx] * (1.0 - weight) + sorted[upper_idx] * weight
112 }
113}
114
115pub fn quantiles(arr: &Array1<f32>, qs: &[f64]) -> Vec<f32> {
126 if arr.is_empty() {
127 return vec![0.0; qs.len()];
128 }
129
130 let mut sorted: Vec<f32> = arr.iter().copied().collect();
131 sorted.sort_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal));
132
133 let n = sorted.len();
134
135 qs.iter()
136 .map(|&q| {
137 let idx_float = q * (n - 1) as f64;
138 let lower_idx = idx_float.floor() as usize;
139 let upper_idx = idx_float.ceil() as usize;
140
141 if lower_idx == upper_idx {
142 sorted[lower_idx]
143 } else {
144 let weight = (idx_float - lower_idx as f64) as f32;
145 sorted[lower_idx] * (1.0 - weight) + sorted[upper_idx] * weight
146 }
147 })
148 .collect()
149}
150
151pub fn normalize_rows(arr: &Array2<f32>) -> Array2<f32> {
161 let mut result = arr.clone();
162 for mut row in result.axis_iter_mut(Axis(0)) {
163 let norm = row.dot(&row).sqrt().max(1e-12);
164 row /= norm;
165 }
166 result
167}
168
169pub fn row_norms(arr: &Array2<f32>) -> Array1<f32> {
179 arr.axis_iter(Axis(0))
180 .map(|row| row.dot(&row).sqrt())
181 .collect()
182}
183
184pub fn packbits(bits: &[u8]) -> Vec<u8> {
194 bits.chunks(8)
195 .map(|chunk| {
196 let mut byte = 0u8;
197 for (i, &bit) in chunk.iter().enumerate() {
198 byte |= bit << (7 - i);
199 }
200 byte
201 })
202 .collect()
203}
204
205pub fn unpackbits(bytes: &[u8]) -> Vec<u8> {
215 let mut bits = Vec::with_capacity(bytes.len() * 8);
216 for &byte in bytes {
217 for i in (0..8).rev() {
218 bits.push((byte >> i) & 1);
219 }
220 }
221 bits
222}
223
224pub fn create_mask(lengths: &ArrayView1<i64>, max_len: usize) -> Array2<bool> {
235 let batch_size = lengths.len();
236 let mut mask = Array2::from_elem((batch_size, max_len), false);
237
238 for (i, &len) in lengths.iter().enumerate() {
239 for j in 0..(len as usize).min(max_len) {
240 mask[[i, j]] = true;
241 }
242 }
243
244 mask
245}
246
247pub fn pad_sequences(sequences: &[Array2<f32>], pad_value: f32) -> (Array2<f32>, Array1<i64>) {
258 if sequences.is_empty() {
259 return (Array2::zeros((0, 0)), Array1::zeros(0));
260 }
261
262 let max_len = sequences.iter().map(|s| s.nrows()).max().unwrap_or(0);
263 let dim = sequences[0].ncols();
264 let batch_size = sequences.len();
265
266 let mut padded = Array2::from_elem((batch_size * max_len, dim), pad_value);
267 let mut lengths = Array1::<i64>::zeros(batch_size);
268
269 for (i, seq) in sequences.iter().enumerate() {
270 let len = seq.nrows();
271 lengths[i] = len as i64;
272 for j in 0..len {
273 for k in 0..dim {
274 padded[[i * max_len + j, k]] = seq[[j, k]];
275 }
276 }
277 }
278
279 (padded, lengths)
280}
281
282#[cfg(test)]
283mod tests {
284 use super::*;
285 use crate::error::Error;
286 use std::io::Write;
287
288 #[test]
289 fn test_quantile() {
290 let arr = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0, 5.0]);
291 assert!((quantile(&arr, 0.5) - 3.0).abs() < 1e-6);
292 assert!((quantile(&arr, 0.0) - 1.0).abs() < 1e-6);
293 assert!((quantile(&arr, 1.0) - 5.0).abs() < 1e-6);
294 }
295
296 #[test]
297 fn test_packbits_unpackbits() {
298 let bits = vec![1, 0, 1, 0, 1, 0, 1, 0, 1, 1, 1, 1, 0, 0, 0, 0];
299 let packed = packbits(&bits);
300 assert_eq!(packed, vec![0b10101010, 0b11110000]);
301
302 let unpacked = unpackbits(&packed);
303 assert_eq!(unpacked, bits);
304 }
305
306 #[test]
307 fn test_normalize_rows() {
308 let arr = Array2::from_shape_vec((2, 3), vec![3.0, 0.0, 4.0, 0.0, 5.0, 0.0]).unwrap();
309 let normalized = normalize_rows(&arr);
310
311 assert!((normalized[[0, 0]] - 0.6).abs() < 1e-6);
313 assert!((normalized[[0, 2]] - 0.8).abs() < 1e-6);
314
315 assert!((normalized[[1, 1]] - 1.0).abs() < 1e-6);
317 }
318
319 #[test]
320 fn atomic_write_failure_preserves_original_file() {
321 let dir = tempfile::tempdir().unwrap();
322 let path = dir.path().join("centroids.npy");
323 std::fs::write(&path, b"original").unwrap();
324
325 let result = atomic_write_file(&path, |file| {
326 file.write_all(b"partial")?;
327 Err(Error::Update("forced failure".to_string()))
328 });
329
330 assert!(result.is_err());
331 assert_eq!(std::fs::read(&path).unwrap(), b"original");
332 let temp_entries: Vec<_> = std::fs::read_dir(dir.path())
333 .unwrap()
334 .filter_map(|entry| entry.ok())
335 .filter(|entry| entry.file_name().to_string_lossy().ends_with(".tmp"))
336 .collect();
337 assert!(temp_entries.is_empty());
338 }
339}