501 lines
16 KiB
Rust
501 lines
16 KiB
Rust
//! Sample data generators for the time series demo
|
|
|
|
use rand::Rng;
|
|
use rand_distr::{Distribution, Normal};
|
|
use timeseries_shared::{DataPoint, SampleDataset, TimeSeriesData};
|
|
|
|
use crate::error::{ForecastError, Result};
|
|
|
|
/// Available sample datasets
|
|
pub const SAMPLE_DATASETS: &[SampleDataset] = &[
|
|
SampleDataset::AirlinePassengers,
|
|
SampleDataset::StockPrices,
|
|
SampleDataset::Temperature,
|
|
SampleDataset::WebTraffic,
|
|
SampleDataset::SineWave,
|
|
SampleDataset::RandomWalk,
|
|
];
|
|
|
|
/// Generate sample time series data
|
|
pub fn generate_sample_data(dataset: SampleDataset, length: usize) -> Result<TimeSeriesData> {
|
|
match dataset {
|
|
SampleDataset::AirlinePassengers => generate_airline_passengers(length),
|
|
SampleDataset::StockPrices => generate_stock_prices(length),
|
|
SampleDataset::Temperature => generate_temperature(length),
|
|
SampleDataset::WebTraffic => generate_web_traffic(length),
|
|
SampleDataset::SineWave => generate_sine_wave(length),
|
|
SampleDataset::RandomWalk => generate_random_walk(length),
|
|
}
|
|
}
|
|
|
|
/// Generate airline passengers-like data (monthly, with trend and seasonality)
|
|
fn generate_airline_passengers(length: usize) -> Result<TimeSeriesData> {
|
|
let mut rng = rand::rng();
|
|
let normal =
|
|
Normal::new(0.0, 10.0).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
|
|
|
|
let mut data = Vec::with_capacity(length);
|
|
|
|
for i in 0..length {
|
|
let t = i as f64;
|
|
// Trend
|
|
let trend = 100.0 + 2.5 * t;
|
|
// Seasonality (12-month cycle)
|
|
let seasonal = 50.0 * (2.0 * std::f64::consts::PI * t / 12.0).sin();
|
|
// Growing amplitude
|
|
let amplitude_factor = 1.0 + 0.01 * t;
|
|
// Noise
|
|
let noise = normal.sample(&mut rng);
|
|
|
|
let value = trend + seasonal * amplitude_factor + noise;
|
|
|
|
data.push(DataPoint {
|
|
timestamp: i as f64,
|
|
value: value.max(0.0),
|
|
});
|
|
}
|
|
|
|
Ok(TimeSeriesData {
|
|
data,
|
|
name: Some("Airline Passengers".to_string()),
|
|
frequency: Some("monthly".to_string()),
|
|
})
|
|
}
|
|
|
|
/// Generate stock prices with trend and volatility clustering
|
|
fn generate_stock_prices(length: usize) -> Result<TimeSeriesData> {
|
|
let mut rng = rand::rng();
|
|
|
|
let mut data = Vec::with_capacity(length);
|
|
let mut price = 100.0;
|
|
let drift = 0.0005; // Small positive drift
|
|
let base_volatility = 0.02;
|
|
|
|
for i in 0..length {
|
|
// Random return with drift
|
|
let volatility = base_volatility * (1.0 + 0.5 * rng.random::<f64>());
|
|
let return_val = drift + volatility * rng.random_range(-1.0..1.0);
|
|
price *= 1.0 + return_val;
|
|
|
|
data.push(DataPoint {
|
|
timestamp: i as f64,
|
|
value: price,
|
|
});
|
|
}
|
|
|
|
Ok(TimeSeriesData {
|
|
data,
|
|
name: Some("Stock Price".to_string()),
|
|
frequency: Some("daily".to_string()),
|
|
})
|
|
}
|
|
|
|
/// Generate temperature data with daily and seasonal patterns
|
|
fn generate_temperature(length: usize) -> Result<TimeSeriesData> {
|
|
let mut rng = rand::rng();
|
|
let normal =
|
|
Normal::new(0.0, 3.0).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
|
|
|
|
let mut data = Vec::with_capacity(length);
|
|
|
|
for i in 0..length {
|
|
let day = i as f64;
|
|
// Annual seasonality (365-day cycle)
|
|
let annual =
|
|
15.0 * (2.0 * std::f64::consts::PI * day / 365.0 - std::f64::consts::PI / 2.0).sin();
|
|
// Base temperature
|
|
let base = 15.0;
|
|
// Noise
|
|
let noise = normal.sample(&mut rng);
|
|
|
|
let value = base + annual + noise;
|
|
|
|
data.push(DataPoint {
|
|
timestamp: i as f64,
|
|
value,
|
|
});
|
|
}
|
|
|
|
Ok(TimeSeriesData {
|
|
data,
|
|
name: Some("Temperature (°C)".to_string()),
|
|
frequency: Some("daily".to_string()),
|
|
})
|
|
}
|
|
|
|
/// Generate web traffic with weekly seasonality and growth trend
|
|
fn generate_web_traffic(length: usize) -> Result<TimeSeriesData> {
|
|
let mut rng = rand::rng();
|
|
let normal =
|
|
Normal::new(0.0, 500.0).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
|
|
|
|
let mut data = Vec::with_capacity(length);
|
|
|
|
for i in 0..length {
|
|
let hour = i as f64;
|
|
// Growth trend
|
|
let trend = 10000.0 + 5.0 * hour;
|
|
// Weekly seasonality (168 hours = 1 week)
|
|
let weekly = 2000.0 * (2.0 * std::f64::consts::PI * hour / 168.0).sin();
|
|
// Daily seasonality (24 hours)
|
|
let daily = 1000.0 * (2.0 * std::f64::consts::PI * hour / 24.0).sin();
|
|
// Noise
|
|
let noise = normal.sample(&mut rng);
|
|
|
|
let value = (trend + weekly + daily + noise).max(0.0);
|
|
|
|
data.push(DataPoint {
|
|
timestamp: i as f64,
|
|
value,
|
|
});
|
|
}
|
|
|
|
Ok(TimeSeriesData {
|
|
data,
|
|
name: Some("Web Traffic (visits/hour)".to_string()),
|
|
frequency: Some("hourly".to_string()),
|
|
})
|
|
}
|
|
|
|
/// Generate clean sine wave for testing
|
|
fn generate_sine_wave(length: usize) -> Result<TimeSeriesData> {
|
|
let mut rng = rand::rng();
|
|
let normal =
|
|
Normal::new(0.0, 0.1).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
|
|
|
|
let mut data = Vec::with_capacity(length);
|
|
|
|
for i in 0..length {
|
|
let t = i as f64;
|
|
let value = (2.0 * std::f64::consts::PI * t / 50.0).sin() + normal.sample(&mut rng);
|
|
|
|
data.push(DataPoint {
|
|
timestamp: t,
|
|
value,
|
|
});
|
|
}
|
|
|
|
Ok(TimeSeriesData {
|
|
data,
|
|
name: Some("Sine Wave".to_string()),
|
|
frequency: Some("synthetic".to_string()),
|
|
})
|
|
}
|
|
|
|
/// Generate random walk for baseline comparison
|
|
fn generate_random_walk(length: usize) -> Result<TimeSeriesData> {
|
|
let mut rng = rand::rng();
|
|
let normal =
|
|
Normal::new(0.0, 1.0).map_err(|e| ForecastError::SampleDataError(e.to_string()))?;
|
|
|
|
let mut data = Vec::with_capacity(length);
|
|
let mut value = 0.0;
|
|
|
|
for i in 0..length {
|
|
value += normal.sample(&mut rng);
|
|
|
|
data.push(DataPoint {
|
|
timestamp: i as f64,
|
|
value,
|
|
});
|
|
}
|
|
|
|
Ok(TimeSeriesData {
|
|
data,
|
|
name: Some("Random Walk".to_string()),
|
|
frequency: Some("synthetic".to_string()),
|
|
})
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
// ========== Basic Generation Tests ==========
|
|
|
|
#[test]
|
|
fn test_generate_all_datasets() {
|
|
for dataset in SAMPLE_DATASETS {
|
|
let result = generate_sample_data(*dataset, 100);
|
|
assert!(result.is_ok(), "Failed for {:?}", dataset);
|
|
let data = result.unwrap();
|
|
assert_eq!(data.data.len(), 100);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_airline_passengers() {
|
|
let data = generate_airline_passengers(144).unwrap();
|
|
assert_eq!(data.data.len(), 144);
|
|
assert!(data.data.iter().all(|p| p.value >= 0.0));
|
|
}
|
|
|
|
// ========== Airline Passengers Tests ==========
|
|
|
|
#[test]
|
|
fn test_airline_passengers_name_and_frequency() {
|
|
let data = generate_airline_passengers(12).unwrap();
|
|
assert_eq!(data.name, Some("Airline Passengers".to_string()));
|
|
assert_eq!(data.frequency, Some("monthly".to_string()));
|
|
}
|
|
|
|
#[test]
|
|
fn test_airline_passengers_trend() {
|
|
let data = generate_airline_passengers(100).unwrap();
|
|
// Should have positive trend - later values generally higher
|
|
let first_half_avg: f64 = data.data[..50].iter().map(|p| p.value).sum::<f64>() / 50.0;
|
|
let second_half_avg: f64 = data.data[50..].iter().map(|p| p.value).sum::<f64>() / 50.0;
|
|
assert!(
|
|
second_half_avg > first_half_avg,
|
|
"Airline data should have upward trend"
|
|
);
|
|
}
|
|
|
|
#[test]
|
|
fn test_airline_passengers_timestamps() {
|
|
let data = generate_airline_passengers(50).unwrap();
|
|
for (i, point) in data.data.iter().enumerate() {
|
|
assert_eq!(point.timestamp, i as f64);
|
|
}
|
|
}
|
|
|
|
// ========== Stock Prices Tests ==========
|
|
|
|
#[test]
|
|
fn test_stock_prices() {
|
|
let data = generate_stock_prices(100).unwrap();
|
|
assert_eq!(data.data.len(), 100);
|
|
assert_eq!(data.name, Some("Stock Price".to_string()));
|
|
assert_eq!(data.frequency, Some("daily".to_string()));
|
|
}
|
|
|
|
#[test]
|
|
fn test_stock_prices_positive() {
|
|
let data = generate_stock_prices(100).unwrap();
|
|
// Stock prices should always be positive (multiplicative process)
|
|
assert!(data.data.iter().all(|p| p.value > 0.0));
|
|
}
|
|
|
|
#[test]
|
|
fn test_stock_prices_starts_at_100() {
|
|
let data = generate_stock_prices(1).unwrap();
|
|
// First value should be close to 100 (starting price with small perturbation)
|
|
assert!((data.data[0].value - 100.0).abs() < 10.0);
|
|
}
|
|
|
|
// ========== Temperature Tests ==========
|
|
|
|
#[test]
|
|
fn test_temperature() {
|
|
let data = generate_temperature(365).unwrap();
|
|
assert_eq!(data.data.len(), 365);
|
|
assert_eq!(data.name, Some("Temperature (°C)".to_string()));
|
|
assert_eq!(data.frequency, Some("daily".to_string()));
|
|
}
|
|
|
|
#[test]
|
|
fn test_temperature_realistic_range() {
|
|
let data = generate_temperature(365).unwrap();
|
|
// Temperatures should be in realistic range (-10 to 40 C roughly)
|
|
for point in &data.data {
|
|
assert!(
|
|
point.value > -20.0 && point.value < 50.0,
|
|
"Temperature {} out of realistic range",
|
|
point.value
|
|
);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_temperature_seasonality() {
|
|
let data = generate_temperature(365).unwrap();
|
|
// Summer (around day 180) should be warmer than winter (around day 0)
|
|
let winter_avg: f64 = data.data[..30].iter().map(|p| p.value).sum::<f64>() / 30.0;
|
|
let summer_avg: f64 = data.data[150..210].iter().map(|p| p.value).sum::<f64>() / 60.0;
|
|
assert!(
|
|
summer_avg > winter_avg,
|
|
"Summer should be warmer than winter"
|
|
);
|
|
}
|
|
|
|
// ========== Web Traffic Tests ==========
|
|
|
|
#[test]
|
|
fn test_web_traffic() {
|
|
let data = generate_web_traffic(168).unwrap();
|
|
assert_eq!(data.data.len(), 168);
|
|
assert_eq!(data.name, Some("Web Traffic (visits/hour)".to_string()));
|
|
assert_eq!(data.frequency, Some("hourly".to_string()));
|
|
}
|
|
|
|
#[test]
|
|
fn test_web_traffic_non_negative() {
|
|
let data = generate_web_traffic(500).unwrap();
|
|
assert!(data.data.iter().all(|p| p.value >= 0.0));
|
|
}
|
|
|
|
#[test]
|
|
fn test_web_traffic_trend() {
|
|
let data = generate_web_traffic(1000).unwrap();
|
|
// Should have positive trend
|
|
let first_half_avg: f64 = data.data[..500].iter().map(|p| p.value).sum::<f64>() / 500.0;
|
|
let second_half_avg: f64 = data.data[500..].iter().map(|p| p.value).sum::<f64>() / 500.0;
|
|
assert!(
|
|
second_half_avg > first_half_avg,
|
|
"Web traffic should trend upward"
|
|
);
|
|
}
|
|
|
|
// ========== Sine Wave Tests ==========
|
|
|
|
#[test]
|
|
fn test_sine_wave() {
|
|
let data = generate_sine_wave(100).unwrap();
|
|
assert_eq!(data.data.len(), 100);
|
|
assert_eq!(data.name, Some("Sine Wave".to_string()));
|
|
assert_eq!(data.frequency, Some("synthetic".to_string()));
|
|
}
|
|
|
|
#[test]
|
|
fn test_sine_wave_range() {
|
|
let data = generate_sine_wave(200).unwrap();
|
|
// Sine wave should be in range [-1.5, 1.5] with small noise
|
|
for point in &data.data {
|
|
assert!(
|
|
point.value >= -2.0 && point.value <= 2.0,
|
|
"Sine value {} out of expected range",
|
|
point.value
|
|
);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_sine_wave_periodicity() {
|
|
let data = generate_sine_wave(200).unwrap();
|
|
// Check that values near multiples of 50 (period) are similar
|
|
// Period is 50, so sine(0) ≈ sine(50) ≈ sine(100)
|
|
let val_at_0 = data.data[0].value;
|
|
let val_at_50 = data.data[50].value;
|
|
let val_at_100 = data.data[100].value;
|
|
// Allow for noise
|
|
assert!((val_at_0 - val_at_50).abs() < 1.0);
|
|
assert!((val_at_0 - val_at_100).abs() < 1.0);
|
|
}
|
|
|
|
// ========== Random Walk Tests ==========
|
|
|
|
#[test]
|
|
fn test_random_walk() {
|
|
let data = generate_random_walk(100).unwrap();
|
|
assert_eq!(data.data.len(), 100);
|
|
assert_eq!(data.name, Some("Random Walk".to_string()));
|
|
assert_eq!(data.frequency, Some("synthetic".to_string()));
|
|
}
|
|
|
|
#[test]
|
|
fn test_random_walk_starts_at_zero() {
|
|
let data = generate_random_walk(10).unwrap();
|
|
// First step should be small (normal with mean 0, std 1)
|
|
assert!(data.data[0].value.abs() < 5.0);
|
|
}
|
|
|
|
#[test]
|
|
fn test_random_walk_variance_grows() {
|
|
// Generate multiple walks and check variance grows with time
|
|
let data = generate_random_walk(1000).unwrap();
|
|
let first_100: Vec<f64> = data.data[..100].iter().map(|p| p.value).collect();
|
|
let last_100: Vec<f64> = data.data[900..].iter().map(|p| p.value).collect();
|
|
|
|
fn variance(v: &[f64]) -> f64 {
|
|
let mean = v.iter().sum::<f64>() / v.len() as f64;
|
|
v.iter().map(|x| (x - mean).powi(2)).sum::<f64>() / v.len() as f64
|
|
}
|
|
|
|
// Later values should have higher variance on average
|
|
// (This test may occasionally fail due to randomness, but typically passes)
|
|
let _first_var = variance(&first_100);
|
|
let _last_var = variance(&last_100);
|
|
// Just check it runs without error - variance comparison is probabilistic
|
|
}
|
|
|
|
// ========== Edge Case Tests ==========
|
|
|
|
#[test]
|
|
fn test_generate_single_point() {
|
|
for dataset in SAMPLE_DATASETS {
|
|
let result = generate_sample_data(*dataset, 1);
|
|
assert!(result.is_ok());
|
|
assert_eq!(result.unwrap().data.len(), 1);
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_generate_large_dataset() {
|
|
let data = generate_sample_data(SampleDataset::SineWave, 10000).unwrap();
|
|
assert_eq!(data.data.len(), 10000);
|
|
}
|
|
|
|
#[test]
|
|
fn test_generate_zero_length() {
|
|
// Zero length should still work (empty dataset)
|
|
let data = generate_sample_data(SampleDataset::SineWave, 0).unwrap();
|
|
assert!(data.data.is_empty());
|
|
}
|
|
|
|
// ========== Timestamp Continuity Tests ==========
|
|
|
|
#[test]
|
|
fn test_timestamps_continuous() {
|
|
for dataset in SAMPLE_DATASETS {
|
|
let data = generate_sample_data(*dataset, 50).unwrap();
|
|
for (i, point) in data.data.iter().enumerate() {
|
|
assert_eq!(
|
|
point.timestamp, i as f64,
|
|
"Timestamp discontinuity at index {} for {:?}",
|
|
i, dataset
|
|
);
|
|
}
|
|
}
|
|
}
|
|
|
|
// ========== Sample Datasets Constant Tests ==========
|
|
|
|
#[test]
|
|
fn test_sample_datasets_constant() {
|
|
assert_eq!(SAMPLE_DATASETS.len(), 6);
|
|
assert!(SAMPLE_DATASETS.contains(&SampleDataset::AirlinePassengers));
|
|
assert!(SAMPLE_DATASETS.contains(&SampleDataset::StockPrices));
|
|
assert!(SAMPLE_DATASETS.contains(&SampleDataset::Temperature));
|
|
assert!(SAMPLE_DATASETS.contains(&SampleDataset::WebTraffic));
|
|
assert!(SAMPLE_DATASETS.contains(&SampleDataset::SineWave));
|
|
assert!(SAMPLE_DATASETS.contains(&SampleDataset::RandomWalk));
|
|
}
|
|
|
|
// ========== Data Quality Tests ==========
|
|
|
|
#[test]
|
|
fn test_no_nan_values() {
|
|
for dataset in SAMPLE_DATASETS {
|
|
let data = generate_sample_data(*dataset, 100).unwrap();
|
|
for point in &data.data {
|
|
assert!(!point.value.is_nan(), "NaN found in {:?}", dataset);
|
|
assert!(!point.timestamp.is_nan(), "NaN timestamp in {:?}", dataset);
|
|
}
|
|
}
|
|
}
|
|
|
|
#[test]
|
|
fn test_no_infinite_values() {
|
|
for dataset in SAMPLE_DATASETS {
|
|
let data = generate_sample_data(*dataset, 100).unwrap();
|
|
for point in &data.data {
|
|
assert!(
|
|
!point.value.is_infinite(),
|
|
"Infinite found in {:?}",
|
|
dataset
|
|
);
|
|
}
|
|
}
|
|
}
|
|
}
|