#![cfg(feature = "serde")]
use std::fs;
use std::path::PathBuf;
use rill_ml::RillError;
use rill_ml::drift::FixedWindowBuffer;
use rill_ml::persistence::{MAX_SNAPSHOT_JSON_BYTES, Snapshot, ValidateState};
use rill_ml::{
bandit::{EpsilonGreedy, LinUcb, ThompsonSampling, Ucb1},
feature_hasher::FeatureHasher,
loss::RegressionLoss,
models::{
BernoulliNaiveBayes, FtrlClassifier, FtrlRegressor, GaussianNaiveBayes, LinearRegression,
MeanRegressor, MultinomialNaiveBayes,
},
optim::{Optimizer, Sgd},
preprocessing::{
ConstantImputer, ForwardFill, FrequencyEncoder, MeanImputer, MissingIndicator,
OneHotEncoder, OrdinalEncoder, StandardScaler,
},
sparse::SparseFeatures,
stats::{ExponentiallyWeightedMean, Mean, Variance},
};
fn fixture_dir(version: &str) -> PathBuf {
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("fixtures")
.join("state")
.join(version)
}
fn load_fixture<T>(version: &str, name: &str) -> T
where
T: serde::de::DeserializeOwned + ValidateState,
{
let path = fixture_dir(version).join(format!("{name}.json"));
let json = fs::read_to_string(&path)
.unwrap_or_else(|e| panic!("read fixture {} {}: {e}", version, path.display()));
Snapshot::<T>::from_json_validated(&json)
.unwrap_or_else(|e| panic!("load fixture {} {}: {e:?}", version, path.display()))
}
fn load_raw(version: &str, name: &str) -> String {
let path = fixture_dir(version).join(format!("{name}.json"));
fs::read_to_string(&path).expect("read fixture")
}
#[test]
fn load_v0_13_0_mean() {
let _: Mean = load_fixture("v0.13.0", "mean");
}
#[test]
fn load_v0_13_0_variance() {
let _: Variance = load_fixture("v0.13.0", "variance");
}
#[test]
fn load_v0_13_0_ew_mean() {
let _: ExponentiallyWeightedMean = load_fixture("v0.13.0", "ew_mean");
}
#[test]
fn load_v0_13_0_standard_scaler() {
let _: StandardScaler = load_fixture("v0.13.0", "standard_scaler");
}
#[test]
fn load_v0_13_0_one_hot_encoder() {
let _: OneHotEncoder = load_fixture("v0.13.0", "one_hot_encoder");
}
#[test]
fn load_v0_13_0_ordinal_encoder() {
let _: OrdinalEncoder = load_fixture("v0.13.0", "ordinal_encoder");
}
#[test]
fn load_v0_13_0_frequency_encoder() {
let _: FrequencyEncoder = load_fixture("v0.13.0", "frequency_encoder");
}
#[test]
fn load_v0_13_0_constant_imputer() {
let _: ConstantImputer = load_fixture("v0.13.0", "constant_imputer");
}
#[test]
fn load_v0_13_0_mean_imputer() {
let _: MeanImputer = load_fixture("v0.13.0", "mean_imputer");
}
#[test]
fn load_v0_13_0_forward_fill() {
let _: ForwardFill = load_fixture("v0.13.0", "forward_fill");
}
#[test]
fn load_v0_13_0_missing_indicator() {
let _: MissingIndicator = load_fixture("v0.13.0", "missing_indicator");
}
#[test]
fn load_v0_13_0_linear_regression() {
let _: LinearRegression = load_fixture("v0.13.0", "linear_regression");
}
#[test]
fn load_v0_13_0_mean_regressor() {
let _: MeanRegressor = load_fixture("v0.13.0", "mean_regressor");
}
#[test]
fn load_v0_13_0_ftrl_regressor() {
let _: FtrlRegressor = load_fixture("v0.13.0", "ftrl_regressor");
}
#[test]
fn load_v0_13_0_ftrl_classifier() {
let _: FtrlClassifier = load_fixture("v0.13.0", "ftrl_classifier");
}
#[test]
fn load_v0_13_0_gaussian_naive_bayes() {
let _: GaussianNaiveBayes = load_fixture("v0.13.0", "gaussian_naive_bayes");
}
#[test]
fn load_v0_13_0_bernoulli_naive_bayes() {
let _: BernoulliNaiveBayes = load_fixture("v0.13.0", "bernoulli_naive_bayes");
}
#[test]
fn load_v0_13_0_multinomial_naive_bayes() {
let _: MultinomialNaiveBayes = load_fixture("v0.13.0", "multinomial_naive_bayes");
}
#[test]
fn load_v0_13_0_optimizer_sgd() {
let json = load_raw("v0.13.0", "optimizer_sgd");
let opt: Optimizer = serde_json::from_str(&json).expect("deserialize optimizer");
let _ = opt;
}
#[test]
fn load_v0_13_0_regression_loss() {
let json = load_raw("v0.13.0", "regression_loss");
let loss: RegressionLoss = serde_json::from_str(&json).expect("deserialize loss");
let _ = loss;
}
#[test]
fn load_v0_13_0_sparse_features() {
let json = load_raw("v0.13.0", "sparse_features");
let _: SparseFeatures = serde_json::from_str(&json).expect("deserialize sparse");
}
#[test]
fn load_v0_13_0_feature_hasher() {
let json = load_raw("v0.13.0", "feature_hasher");
let _: FeatureHasher = serde_json::from_str(&json).expect("deserialize hasher");
}
#[test]
fn load_v0_13_0_epsilon_greedy() {
let _: EpsilonGreedy = load_fixture("v0.13.0", "epsilon_greedy");
}
#[test]
fn load_v0_13_0_ucb1() {
let _: Ucb1 = load_fixture("v0.13.0", "ucb1");
}
#[test]
fn load_v0_13_0_thompson_sampling() {
let _: ThompsonSampling = load_fixture("v0.13.0", "thompson_sampling");
}
#[test]
fn load_v0_13_0_linucb() {
let _: LinUcb = load_fixture("v0.13.0", "linucb");
}
#[test]
fn load_v1_mean() {
let _: Mean = load_fixture("v1", "mean");
}
#[test]
fn load_v1_variance() {
let _: Variance = load_fixture("v1", "variance");
}
#[test]
fn load_v1_ew_mean() {
let _: ExponentiallyWeightedMean = load_fixture("v1", "ew_mean");
}
#[test]
fn load_v1_standard_scaler() {
let _: StandardScaler = load_fixture("v1", "standard_scaler");
}
#[test]
fn load_v1_one_hot_encoder() {
let _: OneHotEncoder = load_fixture("v1", "one_hot_encoder");
}
#[test]
fn load_v1_ordinal_encoder() {
let _: OrdinalEncoder = load_fixture("v1", "ordinal_encoder");
}
#[test]
fn load_v1_frequency_encoder() {
let _: FrequencyEncoder = load_fixture("v1", "frequency_encoder");
}
#[test]
fn load_v1_constant_imputer() {
let _: ConstantImputer = load_fixture("v1", "constant_imputer");
}
#[test]
fn load_v1_mean_imputer() {
let _: MeanImputer = load_fixture("v1", "mean_imputer");
}
#[test]
fn load_v1_forward_fill() {
let _: ForwardFill = load_fixture("v1", "forward_fill");
}
#[test]
fn load_v1_missing_indicator() {
let _: MissingIndicator = load_fixture("v1", "missing_indicator");
}
#[test]
fn load_v1_linear_regression() {
let _: LinearRegression = load_fixture("v1", "linear_regression");
}
#[test]
fn load_v1_mean_regressor() {
let _: MeanRegressor = load_fixture("v1", "mean_regressor");
}
#[test]
fn load_v1_ftrl_regressor() {
let _: FtrlRegressor = load_fixture("v1", "ftrl_regressor");
}
#[test]
fn load_v1_ftrl_classifier() {
let _: FtrlClassifier = load_fixture("v1", "ftrl_classifier");
}
#[test]
fn load_v1_gaussian_naive_bayes() {
let _: GaussianNaiveBayes = load_fixture("v1", "gaussian_naive_bayes");
}
#[test]
fn load_v1_bernoulli_naive_bayes() {
let _: BernoulliNaiveBayes = load_fixture("v1", "bernoulli_naive_bayes");
}
#[test]
fn load_v1_multinomial_naive_bayes() {
let _: MultinomialNaiveBayes = load_fixture("v1", "multinomial_naive_bayes");
}
#[test]
fn load_v1_optimizer_sgd() {
let json = load_raw("v1", "optimizer_sgd");
let _: Optimizer = serde_json::from_str(&json).expect("deserialize optimizer");
}
#[test]
fn load_v1_regression_loss() {
let json = load_raw("v1", "regression_loss");
let _: RegressionLoss = serde_json::from_str(&json).expect("deserialize loss");
}
#[test]
fn load_v1_sparse_features() {
let json = load_raw("v1", "sparse_features");
let _: SparseFeatures = serde_json::from_str(&json).expect("deserialize sparse");
}
#[test]
fn load_v1_feature_hasher() {
let json = load_raw("v1", "feature_hasher");
let _: FeatureHasher = serde_json::from_str(&json).expect("deserialize hasher");
}
#[test]
fn load_v1_epsilon_greedy() {
let _: EpsilonGreedy = load_fixture("v1", "epsilon_greedy");
}
#[test]
fn load_v1_ucb1() {
let _: Ucb1 = load_fixture("v1", "ucb1");
}
#[test]
fn load_v1_thompson_sampling() {
let _: ThompsonSampling = load_fixture("v1", "thompson_sampling");
}
#[test]
fn load_v1_linucb() {
let _: LinUcb = load_fixture("v1", "linucb");
}
#[test]
fn incompatible_format_version_rejected() {
let json = r#"{"format_version":999,"model":{"count":1,"mean":1.0}}"#;
let result: Result<Mean, _> = Snapshot::from_json_validated(json);
assert!(result.is_err());
}
#[test]
fn null_mean_rejected_by_serde() {
let json = r#"{"format_version":1,"model":{"count":1,"mean":null}}"#;
let result: Result<Mean, _> = Snapshot::from_json_validated(json);
assert!(result.is_err());
}
#[test]
fn validate_state_rejects_non_finite_values() {
#[derive(serde::Serialize, serde::Deserialize)]
struct FloatWrapper {
raw: String,
}
impl ValidateState for FloatWrapper {
fn validate_state(&self) -> Result<(), RillError> {
let value: f64 = self.raw.parse().map_err(|_| {
RillError::InvalidState(format!("invalid float literal: {}", self.raw))
})?;
if !value.is_finite() {
return Err(RillError::NonFiniteValue {
field: "value",
value,
});
}
Ok(())
}
}
let json = r#"{"format_version":1,"model":{"raw":"NaN"}}"#;
let result: Result<FloatWrapper, _> = Snapshot::from_json_validated(json);
assert!(result.is_err(), "ValidateState must reject NaN");
let json = r#"{"format_version":1,"model":{"raw":"inf"}}"#;
let result: Result<FloatWrapper, _> = Snapshot::from_json_validated(json);
assert!(result.is_err(), "ValidateState must reject Infinity");
let json = r#"{"format_version":1,"model":{"raw":"-inf"}}"#;
let result: Result<FloatWrapper, _> = Snapshot::from_json_validated(json);
assert!(result.is_err(), "ValidateState must reject -Infinity");
let json = r#"{"format_version":1,"model":{"raw":"1.5"}}"#;
let result: Result<FloatWrapper, _> = Snapshot::from_json_validated(json);
assert!(result.is_ok(), "ValidateState must accept finite values");
}
#[test]
fn oversized_json_rejected_before_deserialization() {
let payload = "x".repeat(MAX_SNAPSHOT_JSON_BYTES + 1);
let json = format!(
"{{\"format_version\":1,\"model\":{{\"count\":1,\"mean\":1.0}},\"pad\":\"{payload}\"}}"
);
let result: Result<Mean, _> = Snapshot::from_json_validated(&json);
assert!(result.is_err());
}
#[test]
fn dimension_mismatch_in_linear_regression_rejected() {
let json = r#"{
"format_version":1,
"model":{
"feature_count":3,
"weights":[0.0,0.0],
"intercept":0.0,
"optimizer":{"Sgd":{"feature_count":3,"config":{"learning_rate":0.01,"l2":0.0},"samples_seen":0}},
"loss":{"Mse":{}},
"samples_seen":0
}
}"#;
let result: Result<LinearRegression, _> = Snapshot::from_json_validated(json);
assert!(result.is_err(), "must reject dimension mismatch");
}
#[test]
fn optimizer_param_count_mismatch_rejected() {
let json = r#"{
"format_version":1,
"model":{
"feature_count":2,
"weights":[0.0,0.0],
"intercept":0.0,
"optimizer":{"Sgd":{"feature_count":1,"config":{"learning_rate":0.01,"l2":0.0},"samples_seen":0}},
"loss":{"Mse":{}},
"samples_seen":0
}
}"#;
let result: Result<LinearRegression, _> = Snapshot::from_json_validated(json);
assert!(result.is_err(), "must reject optimizer param mismatch");
}
#[test]
fn failed_restore_returns_no_model() {
let json = r#"{"format_version":999,"model":{"count":1,"mean":1.0}}"#;
let result: Result<Mean, _> = Snapshot::from_json_validated(json);
assert!(result.is_err());
}
#[test]
fn negative_m2_in_variance_rejected() {
let json = r#"{"format_version":1,"model":{"count":1,"mean":1.0,"m2":-1.0,"kind":"Sample"}}"#;
let result: Result<Variance, _> = Snapshot::from_json_validated(json);
assert!(result.is_err(), "must reject negative m2 in Variance state");
}
#[test]
fn sgd_invalid_learning_rate_rejected() {
let json = r#"{"format_version":1,"model":{"feature_count":2,"config":{"learning_rate":0.0,"l2":0.0},"samples_seen":0}}"#;
let result: Result<Sgd, _> = Snapshot::from_json_validated(json);
assert!(
result.is_err(),
"must reject non-positive learning rate in Sgd state"
);
}
#[test]
fn encoder_mapping_inconsistency_rejected() {
let json = r#"{
"format_version":1,
"model":{
"categories":["beta","alpha"],
"samples_seen":2
}
}"#;
let result: Result<OneHotEncoder, _> = Snapshot::from_json_validated(json);
assert!(
result.is_err(),
"must reject unsorted encoder categories because indices would be inconsistent"
);
}
#[test]
fn bandit_arm_count_mismatch_rejected() {
let json = r#"{
"format_version":1,
"model":{
"arm_count":3,
"config":{"epsilon":0.1,"decay":1.0,"min_epsilon":0.0},
"pulls":[1,2],
"total_rewards":[1.0,2.0],
"samples_seen":3,
"current_epsilon":0.1
}
}"#;
let result: Result<EpsilonGreedy, _> = Snapshot::from_json_validated(json);
assert!(
result.is_err(),
"must reject bandit vectors that do not match arm_count"
);
}
#[test]
fn drift_buffer_capacity_overflow_rejected() {
let json = r#"{
"format_version":1,
"model":{
"buffer":[1.0,2.0,3.0],
"capacity":3,
"head":1,
"len":4
}
}"#;
let result: Result<Snapshot<FixedWindowBuffer>, _> = serde_json::from_str(json);
assert!(
result.is_err(),
"must reject a drift buffer whose length exceeds its capacity"
);
}