use core::time::Duration;
use serde::{Deserialize, Serialize};
use crate::height_vector::HeightVector;
#[cfg(feature = "f32")]
type FloatType = f32;
#[cfg(not(feature = "f32"))]
type FloatType = f64;
const C_ERROR: FloatType = 0.25;
const C_DELTA: FloatType = 0.25;
const DEFAULT_ERROR: FloatType = 200.0;
const MIN_ERROR: FloatType = FloatType::EPSILON;
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct NetworkCoordinate<const N: usize> {
#[serde(flatten)]
heightvec: HeightVector<N>,
error: FloatType,
}
#[allow(clippy::module_name_repetitions)]
pub type NetworkCoordinate2D = NetworkCoordinate<2>;
#[allow(clippy::module_name_repetitions)]
pub type NetworkCoordinate3D = NetworkCoordinate<3>;
impl<const N: usize> NetworkCoordinate<N> {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn estimated_rtt(&self, rhs: &Self) -> Duration {
cfg_if::cfg_if! {
if #[cfg(feature = "f32")] {
Duration::from_secs_f32((self.heightvec - rhs.heightvec).len() / 1000.0)
} else {
Duration::from_secs_f64((self.heightvec - rhs.heightvec).len() / 1000.0)
}
}
}
pub fn update(&mut self, rhs: &Self, rtt: Duration) -> &Self {
cfg_if::cfg_if! {
if #[cfg(feature = "f32")] {
let rtt_ms = rtt.as_secs_f32() * 1000.0;
let rtt_estimated_ms = self.estimated_rtt(rhs).as_secs_f32() * 1000.0;
} else {
let rtt_ms = rtt.as_secs_f64() * 1000.0;
let rtt_estimated_ms = self.estimated_rtt(rhs).as_secs_f64() * 1000.0;
}
}
if rtt_ms < 0.0 {
unreachable!();
}
let w = self.error / (self.error + rhs.error);
let error = rtt_ms - rtt_estimated_ms;
let es = error.abs() / rtt_ms;
self.error = (es * C_ERROR)
.mul_add(w, self.error * C_ERROR.mul_add(-w, 1.0))
.max(MIN_ERROR);
let delta = C_DELTA * w;
self.heightvec =
self.heightvec + (self.heightvec - rhs.heightvec).normalized() * delta * error;
if self.heightvec.is_invalid() {
*self = Self::new();
unreachable!();
}
self
}
#[must_use]
pub const fn error(&self) -> FloatType {
self.error
}
}
impl<const N: usize> Default for NetworkCoordinate<N> {
fn default() -> Self {
Self {
heightvec: HeightVector::<N>::random(),
error: DEFAULT_ERROR,
}
}
}
#[cfg(test)]
mod tests {
use assert_approx_eq::assert_approx_eq;
use super::*;
#[test]
fn test_convergence() {
let mut a = NetworkCoordinate::<3>::new();
let mut b = NetworkCoordinate::<3>::new();
let t = Duration::from_millis(250);
(0..20).for_each(|_| {
a.update(&b, t);
b.update(&a, t);
});
let rtt = a.estimated_rtt(&b);
assert_approx_eq!(rtt.as_secs_f32() * 1000.0, 250.0, 1.0);
}
#[test]
fn test_mini_network() {
let mut slc = NetworkCoordinate::<2>::new();
let mut nyc = NetworkCoordinate::<2>::new();
let mut lax = NetworkCoordinate::<2>::new();
let mut mad = NetworkCoordinate::<2>::new();
let error = slc.error.hypot(nyc.error.hypot(lax.error.hypot(mad.error)));
assert_approx_eq!(error, 400.0);
(0..20).for_each(|_| {
slc.update(&nyc, Duration::from_millis(162));
nyc.update(&slc, Duration::from_millis(162));
slc.update(&lax, Duration::from_millis(115));
lax.update(&slc, Duration::from_millis(115));
slc.update(&mad, Duration::from_millis(242));
mad.update(&slc, Duration::from_millis(242));
nyc.update(&lax, Duration::from_millis(95));
lax.update(&nyc, Duration::from_millis(95));
nyc.update(&mad, Duration::from_millis(168));
mad.update(&nyc, Duration::from_millis(168));
lax.update(&mad, Duration::from_millis(192));
mad.update(&lax, Duration::from_millis(192));
});
let error = slc.error + nyc.error + lax.error + mad.error;
assert!(error < 5.0);
}
#[test]
fn test_serde() {
let s = "{\"position\":[1.5,0.5,2.0],\"height\":0.1,\"error\":1.0}";
let a: NetworkCoordinate<3> =
serde_json::from_str(s).expect("deserialization failed during test");
assert_approx_eq!(a.heightvec.len(), 2.649_509, 0.001);
assert_approx_eq!(a.error, 1.0);
assert_eq!(a.estimated_rtt(&a).as_millis(), 0);
let t = serde_json::to_string(&a);
assert_eq!(t.as_ref().expect("serialization failed during test"), s);
}
#[test]
fn test_estimated_rtt() {
let s = "{\"position\":[1.5,0.5,2.0],\"height\":25.0,\"error\":1.0}";
let a: NetworkCoordinate<3> =
serde_json::from_str(s).expect("deserialization failed during test");
let s = "{\"position\":[-1.5,-0.5,-2.0],\"height\":50.0,\"error\":1.0}";
let b: NetworkCoordinate<3> =
serde_json::from_str(s).expect("deserialization failed during test");
let estimate = a.estimated_rtt(&b);
assert_approx_eq!(estimate.as_secs_f32(), 0.080_099);
}
#[test]
fn test_error_getter() {
let s = "{\"position\":[1.5,0.5,2.0],\"height\":25.0,\"error\":1.0}";
let a: NetworkCoordinate<3> =
serde_json::from_str(s).expect("deserialization failed during test");
assert_approx_eq!(a.error(), 1.0);
}
}