use crate::protocol::sd::RebootFlag;
use std::{collections::HashMap, net::SocketAddr};
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
pub enum TransportKind {
Multicast,
#[allow(dead_code)]
Unicast,
}
type SessionKey = (SocketAddr, TransportKind, u16, u16);
#[derive(Clone, Copy, Debug)]
struct SessionState {
last_session_id: u16,
last_reboot_flag: RebootFlag,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum SessionVerdict {
Ok,
Reboot,
Initial,
}
#[derive(Debug, Default)]
pub struct SessionTracker {
state: HashMap<SessionKey, SessionState>,
}
impl SessionTracker {
pub fn check(
&mut self,
sender: SocketAddr,
transport: TransportKind,
service_id: u16,
instance_id: u16,
session_id: u16,
reboot_flag: RebootFlag,
) -> SessionVerdict {
let key = (sender, transport, service_id, instance_id);
let verdict = match self.state.get(&key) {
None => SessionVerdict::Initial,
Some(prev) => {
if prev.last_reboot_flag == RebootFlag::Continuous
&& reboot_flag == RebootFlag::RecentlyRebooted
{
SessionVerdict::Reboot
} else if prev.last_reboot_flag == RebootFlag::RecentlyRebooted
&& reboot_flag == RebootFlag::RecentlyRebooted
&& session_id < prev.last_session_id
&& !(prev.last_session_id == u16::MAX && session_id <= 1)
{
SessionVerdict::Reboot
} else {
SessionVerdict::Ok
}
}
};
self.state.insert(
key,
SessionState {
last_session_id: session_id,
last_reboot_flag: reboot_flag,
},
);
verdict
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::net::{Ipv4Addr, SocketAddr};
fn addr(port: u16) -> SocketAddr {
SocketAddr::new(Ipv4Addr::new(192, 168, 1, 10).into(), port)
}
const SVC: u16 = 0x0047;
const INST: u16 = 0x0001;
const SVC_B: u16 = 0x005D;
const RB: RebootFlag = RebootFlag::RecentlyRebooted;
const CONT: RebootFlag = RebootFlag::Continuous;
#[test]
fn first_message_returns_initial() {
let mut tracker = SessionTracker::default();
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
assert_eq!(verdict, SessionVerdict::Initial);
}
#[test]
fn normal_increment_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 2, RB);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn reboot_flag_continuous_to_recently_rebooted_returns_reboot() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, CONT);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
assert_eq!(verdict, SessionVerdict::Reboot);
}
#[test]
fn session_id_decrease_same_service_with_recently_rebooted_returns_reboot() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, RB);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 50, RB);
assert_eq!(verdict, SessionVerdict::Reboot);
}
#[test]
fn session_id_decrease_different_services_no_false_reboot() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, RB);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 50, RB);
assert_eq!(verdict, SessionVerdict::Initial);
}
#[test]
fn interleaved_sd_offers_no_false_reboot() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 1, RB);
assert_eq!(v, SessionVerdict::Initial);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 2, RB);
assert_eq!(v, SessionVerdict::Ok);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 2, RB);
assert_eq!(v, SessionVerdict::Ok);
}
#[test]
fn session_id_decrease_with_continuous_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, CONT);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 50, CONT);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn different_transports_tracked_separately() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, RB);
let verdict = tracker.check(addr(1000), TransportKind::Unicast, SVC, INST, 1, RB);
assert_eq!(verdict, SessionVerdict::Initial);
}
#[test]
fn different_senders_tracked_separately() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, RB);
let verdict = tracker.check(addr(2000), TransportKind::Multicast, SVC, INST, 1, RB);
assert_eq!(verdict, SessionVerdict::Initial);
}
#[test]
fn reboot_flag_recently_rebooted_to_continuous_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, RB);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 101, CONT);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn same_session_id_with_recently_rebooted_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 5, RB);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 5, RB);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn different_instance_ids_tracked_separately() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, 0x0001, 100, RB);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, 0x0002, 1, RB);
assert_eq!(verdict, SessionVerdict::Initial);
}
#[test]
fn session_id_wrap_around_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 65535, RB);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn session_id_wrap_around_then_normal_increment() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 65535, RB);
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 2, RB);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn session_id_wrap_to_zero_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 65535, RB);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 0, RB);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn reboot_flag_transition_with_session_id_decrease_both_signal_reboot() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, CONT);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
assert_eq!(verdict, SessionVerdict::Reboot);
}
#[test]
fn multiple_reboots_in_sequence() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 50, RB);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
assert_eq!(v, SessionVerdict::Reboot);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 2, RB);
assert_eq!(v, SessionVerdict::Ok);
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 10, CONT);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
assert_eq!(v, SessionVerdict::Reboot);
}
#[test]
fn interleaved_offers_with_real_reboot() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 10, RB);
tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 10, RB);
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 11, RB);
tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 11, RB);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, RB);
assert_eq!(v, SessionVerdict::Reboot);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 1, RB);
assert_eq!(v, SessionVerdict::Reboot);
}
#[test]
fn normal_increment_with_continuous_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, CONT);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 2, CONT);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn interleaved_transports_for_same_instance_do_not_false_reboot() {
let mut t = SessionTracker::default();
let a = addr(30490);
assert_eq!(
t.check(a, TransportKind::Multicast, SVC, INST, 1468, RB),
SessionVerdict::Initial
);
assert_eq!(
t.check(a, TransportKind::Unicast, SVC, INST, 739, RB),
SessionVerdict::Initial
);
assert_eq!(
t.check(a, TransportKind::Multicast, SVC, INST, 1469, RB),
SessionVerdict::Ok
);
assert_eq!(
t.check(a, TransportKind::Unicast, SVC, INST, 740, RB),
SessionVerdict::Ok
);
assert_eq!(
t.check(a, TransportKind::Multicast, SVC, INST, 3, RB),
SessionVerdict::Reboot
);
}
#[test]
fn same_transport_mis_tag_false_reboots() {
let mut t = SessionTracker::default();
let a = addr(30490);
t.check(a, TransportKind::Multicast, SVC, INST, 1468, RB);
assert_eq!(
t.check(a, TransportKind::Multicast, SVC, INST, 739, RB),
SessionVerdict::Reboot
);
}
}