use std::path::Path;
use crate::error::Result;
pub fn to_json<T: serde::Serialize>(t: &T) -> Result<String> {
Ok(serde_json::to_string_pretty(t)?)
}
pub fn from_json<T: serde::de::DeserializeOwned>(s: &str) -> Result<T> {
Ok(serde_json::from_str(s)?)
}
pub fn save_json<T: serde::Serialize, P: AsRef<Path>>(t: &T, path: P) -> Result<()> {
let s = to_json(t)?;
std::fs::write(path, s)?;
Ok(())
}
pub fn load_json<T: serde::de::DeserializeOwned, P: AsRef<Path>>(path: P) -> Result<T> {
let s = std::fs::read_to_string(path)?;
from_json(&s)
}
#[cfg(all(test, feature = "serde"))]
mod tests {
use super::*;
use crate::matrix::Matrix;
use crate::scaler::StandardScaler;
use crate::traits::Transformer;
#[test]
fn to_from_json_round_trips_scaler_params() {
let x = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0], vec![5.0, 6.0]]).unwrap();
let mut scaler = StandardScaler::new();
scaler.fit(&x).unwrap();
let json = to_json(&scaler).unwrap();
let restored: StandardScaler = from_json(&json).unwrap();
assert_eq!(restored.mean(), scaler.mean());
assert_eq!(restored.std(), scaler.std());
}
#[test]
fn save_json_creates_file_and_load_restores() {
let x = Matrix::new(vec![vec![1.0, 2.0], vec![3.0, 4.0], vec![5.0, 6.0]]).unwrap();
let mut scaler = StandardScaler::new();
scaler.fit(&x).unwrap();
let path = std::env::temp_dir().join("datarust_serialize_test.json");
let _ = std::fs::remove_file(&path);
save_json(&scaler, &path).unwrap();
assert!(path.exists(), "save_json did not create the file");
let restored: StandardScaler = load_json(&path).unwrap();
let original_out = scaler.transform(&x).unwrap();
let restored_out = restored.transform(&x).unwrap();
for i in 0..x.nrows() {
for j in 0..x.ncols() {
assert!(
(original_out.get(i, j) - restored_out.get(i, j)).abs() < 1e-12,
"i={i} j={j}"
);
}
}
let _ = std::fs::remove_file(&path);
}
#[test]
fn load_json_missing_file_errors() {
let path = std::env::temp_dir().join("datarust_does_not_exist_12345.json");
let _ = std::fs::remove_file(&path);
let result: Result<StandardScaler> = load_json(&path);
assert!(result.is_err(), "loading a missing file should error");
assert!(matches!(
result.unwrap_err(),
crate::error::DatarustError::Io(_)
));
}
#[test]
fn from_json_malformed_errors() {
let result: Result<StandardScaler> = from_json("not valid json {");
assert!(result.is_err());
assert!(matches!(
result.unwrap_err(),
crate::error::DatarustError::Serde(_)
));
}
}