use differential_equations::traits::State as StateTrait;
use differential_equations_derive::State;
use nalgebra::{Matrix2, Vector2, Vector3};
use num_complex::Complex;
fn test_state_basics<S: StateTrait<f64>>(state: &mut S, expected_len: usize) {
assert_eq!(state.len(), expected_len);
for i in 0..expected_len {
let original = state.get_component(i);
state.set_component(i, 42.0);
assert_eq!(state.get_component(i), 42.0);
state.set_component(i, original); }
state.map_components_mut(|i, val| {
*val = i as f64 + 1.0;
});
for i in 0..expected_len {
assert_eq!(state.get_component(i), i as f64 + 1.0);
}
}
#[test]
fn test_single_field_only() {
#[derive(State, PartialEq)]
struct SingleField<T> {
x: T,
}
let mut state = SingleField { x: 1.0 };
test_state_basics(&mut state, 1);
let other = SingleField { x: 2.0 };
let sum = state + other;
assert_eq!(sum.x, 3.0);
let difference = sum - SingleField { x: 1.0 };
assert_eq!(difference.x, 2.0);
let scaled = difference * 3.0;
assert_eq!(scaled.x, 6.0);
let divided = scaled / 2.0;
assert_eq!(divided.x, 3.0);
let zero = SingleField::<f64>::zeros();
assert_eq!(zero.x, 0.0);
}
#[test]
fn test_multiple_single_fields() {
#[derive(State)]
struct MultipleFields<T> {
x: T,
y: T,
z: T,
}
let mut state = MultipleFields {
x: 1.0,
y: 2.0,
z: 3.0,
};
test_state_basics(&mut state, 3);
assert_eq!(state.get_component(0), 1.0);
assert_eq!(state.get_component(1), 2.0);
assert_eq!(state.get_component(2), 3.0);
let other = MultipleFields {
x: 1.0,
y: 1.0,
z: 1.0,
};
let sum = state + other;
assert_eq!(sum.x, 2.0);
assert_eq!(sum.y, 3.0);
assert_eq!(sum.z, 4.0);
}
#[test]
fn test_non_generic_named_state() {
#[derive(State, PartialEq)]
struct ConcreteState {
x: f64,
y: f64,
velocity: [f64; 2],
}
let mut state = ConcreteState {
x: 1.0,
y: 2.0,
velocity: [3.0, 4.0],
};
test_state_basics(&mut state, 4);
let other = ConcreteState {
x: 0.5,
y: 1.5,
velocity: [2.0, 3.0],
};
let sum = state + other;
assert_eq!(sum.x, 1.5);
assert_eq!(sum.y, 3.5);
assert_eq!(sum.velocity, [5.0, 7.0]);
let scaled = sum * 2.0;
assert_eq!(scaled.x, 3.0);
assert_eq!(scaled.y, 7.0);
assert_eq!(scaled.velocity, [10.0, 14.0]);
let zero = ConcreteState::zeros();
assert_eq!(zero.x, 0.0);
assert_eq!(zero.y, 0.0);
assert_eq!(zero.velocity, [0.0, 0.0]);
}
#[test]
fn test_non_generic_tuple_state() {
#[derive(State, PartialEq)]
struct ConcreteTuple(f64, [f64; 2]);
let mut state = ConcreteTuple(1.0, [2.0, 3.0]);
test_state_basics(&mut state, 3);
let difference = state - ConcreteTuple(1.0, [1.0, 1.0]);
assert_eq!(difference.0, 0.0);
assert_eq!(difference.1, [1.0, 2.0]);
}
#[test]
fn test_array_fields_only() {
#[derive(State)]
struct ArrayFields<T> {
small_array: [T; 2],
large_array: [T; 5],
}
let mut state = ArrayFields {
small_array: [1.0, 2.0],
large_array: [3.0, 4.0, 5.0, 6.0, 7.0],
};
test_state_basics(&mut state, 7);
assert_eq!(state.get_component(0), 1.0); assert_eq!(state.get_component(1), 2.0); assert_eq!(state.get_component(2), 3.0); assert_eq!(state.get_component(6), 7.0);
let other = ArrayFields {
small_array: [0.5, 0.5],
large_array: [1.0, 1.0, 1.0, 1.0, 1.0],
};
let sum = state + other;
assert_eq!(sum.small_array[0], 1.5);
assert_eq!(sum.large_array[2], 6.0);
}
#[test]
fn test_nalgebra_fields_only() {
#[derive(State)]
struct NalgebraFields<T> {
vec2: Vector2<T>,
vec3: Vector3<T>,
mat2: Matrix2<T>,
}
let mut state = NalgebraFields {
vec2: Vector2::new(1.0, 2.0),
vec3: Vector3::new(3.0, 4.0, 5.0),
mat2: Matrix2::new(6.0, 7.0, 8.0, 9.0),
};
test_state_basics(&mut state, 9);
assert_eq!(state.get_component(0), 1.0); assert_eq!(state.get_component(1), 2.0); assert_eq!(state.get_component(2), 3.0); assert_eq!(state.get_component(4), 5.0); assert_eq!(state.get_component(5), 6.0); assert_eq!(state.get_component(6), 7.0); assert_eq!(state.get_component(7), 8.0); assert_eq!(state.get_component(8), 9.0); }
#[test]
fn test_complex_fields_only() {
#[derive(State)]
struct ComplexFields<T> {
z1: Complex<T>,
z2: Complex<T>,
}
let mut state = ComplexFields {
z1: Complex::new(1.0, 2.0),
z2: Complex::new(3.0, 4.0),
};
test_state_basics(&mut state, 4);
assert_eq!(state.get_component(0), 1.0); assert_eq!(state.get_component(1), 2.0); assert_eq!(state.get_component(2), 3.0); assert_eq!(state.get_component(3), 4.0);
let other = ComplexFields {
z1: Complex::new(0.5, 0.5),
z2: Complex::new(1.0, 1.0),
};
let sum = state + other;
assert_eq!(sum.z1.re, 1.5);
assert_eq!(sum.z1.im, 2.5);
assert_eq!(sum.z2.re, 4.0);
assert_eq!(sum.z2.im, 5.0);
}
#[test]
fn test_mixed_field_types() {
#[derive(State)]
struct MixedState<T> {
scalar: T,
array: [T; 3],
vector: Vector2<T>,
complex: Complex<T>,
}
let mut state = MixedState {
scalar: 1.0,
array: [2.0, 3.0, 4.0],
vector: Vector2::new(5.0, 6.0),
complex: Complex::new(7.0, 8.0),
};
test_state_basics(&mut state, 8);
assert_eq!(state.get_component(0), 1.0); assert_eq!(state.get_component(1), 2.0); assert_eq!(state.get_component(3), 4.0); assert_eq!(state.get_component(4), 5.0); assert_eq!(state.get_component(5), 6.0); assert_eq!(state.get_component(6), 7.0); assert_eq!(state.get_component(7), 8.0); }