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: bool,
}
#[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: bool,
) -> 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 && reboot_flag {
SessionVerdict::Reboot
} else if prev.last_reboot_flag && reboot_flag && session_id < prev.last_session_id
{
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;
#[test]
fn first_message_returns_initial() {
let mut tracker = SessionTracker::default();
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, true);
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, true);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 2, true);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn reboot_flag_0_to_1_returns_reboot() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, false);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, true);
assert_eq!(verdict, SessionVerdict::Reboot);
}
#[test]
fn session_id_decrease_same_service_with_reboot_flag_1_returns_reboot() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, true);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 50, true);
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, true);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 50, true);
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, true);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 1, true);
assert_eq!(v, SessionVerdict::Initial);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 2, true);
assert_eq!(v, SessionVerdict::Ok);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 2, true);
assert_eq!(v, SessionVerdict::Ok);
}
#[test]
fn session_id_decrease_with_reboot_flag_0_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, false);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 50, false);
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, true);
let verdict = tracker.check(addr(1000), TransportKind::Unicast, SVC, INST, 1, true);
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, true);
let verdict = tracker.check(addr(2000), TransportKind::Multicast, SVC, INST, 1, true);
assert_eq!(verdict, SessionVerdict::Initial);
}
#[test]
fn reboot_flag_1_to_0_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 100, true);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 101, false);
assert_eq!(verdict, SessionVerdict::Ok);
}
#[test]
fn same_session_id_with_reboot_flag_1_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 5, true);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 5, true);
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, true);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, 0x0002, 1, true);
assert_eq!(verdict, SessionVerdict::Initial);
}
#[test]
fn session_id_wrap_around_currently_treated_as_reboot() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 65535, true);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, true);
assert_eq!(verdict, SessionVerdict::Reboot);
}
#[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, false);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, true);
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, true);
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 50, true);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, true);
assert_eq!(v, SessionVerdict::Reboot);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 2, true);
assert_eq!(v, SessionVerdict::Ok);
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 10, false);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, true);
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, true);
tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 10, true);
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 11, true);
tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 11, true);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, true);
assert_eq!(v, SessionVerdict::Reboot);
let v = tracker.check(addr(1000), TransportKind::Multicast, SVC_B, INST, 1, true);
assert_eq!(v, SessionVerdict::Reboot);
}
#[test]
fn normal_increment_with_reboot_flag_0_returns_ok() {
let mut tracker = SessionTracker::default();
tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 1, false);
let verdict = tracker.check(addr(1000), TransportKind::Multicast, SVC, INST, 2, false);
assert_eq!(verdict, SessionVerdict::Ok);
}
}