use std::sync::mpsc::{self, Receiver, RecvTimeoutError, Sender};
use std::thread;
use std::time::Duration;
use crate::device::Device;
use crate::error::Result;
use crate::snapshot::{Snapshot, merge};
use crate::timestamp::Timestamp;
use crate::{iokit, presence, profiler};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct Tiers {
pub fast: Duration,
pub slow: Duration,
pub timeout: Duration,
}
impl Default for Tiers {
fn default() -> Self {
Self {
fast: Duration::from_secs(30),
slow: Duration::from_secs(300),
timeout: Duration::from_secs(10),
}
}
}
pub fn snapshot() -> Snapshot {
let read_at = Timestamp::now();
let timeout = Tiers::default().timeout;
let cached = read_slow(&Cached::default(), read_at, |at, warnings| {
profiler::read(at, timeout, warnings)
});
read_fast(read_at, iokit::read, &cached)
}
pub fn poll(tiers: Tiers) -> Receiver<Snapshot> {
poll_with(
tiers,
iokit::read,
move |at, warnings| profiler::read(at, tiers.timeout, warnings),
Timestamp::now,
presence::watch(),
)
}
fn poll_with<F, S, C>(
tiers: Tiers,
fast: F,
slow: S,
clock: C,
nudges: Receiver<()>,
) -> Receiver<Snapshot>
where
F: Fn(Timestamp, &mut Vec<String>) -> Vec<Device> + Send + 'static,
S: Fn(Timestamp, &mut Vec<String>) -> Result<Vec<Device>> + Send + 'static,
C: Fn() -> Timestamp + Clone + Send + 'static,
{
let (snapshots, readings) = mpsc::channel();
let (refreshed, cached) = mpsc::channel();
let (polling, wanted) = mpsc::channel();
let slow_clock = clock.clone();
thread::spawn(move || slow_tier(tiers.slow, slow, slow_clock, &refreshed, &wanted));
thread::spawn(move || {
fast_tier(
tiers.fast, fast, clock, &snapshots, &cached, polling, &nudges,
)
});
readings
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
struct Cached {
devices: Vec<Device>,
warnings: Vec<String>,
degraded: bool,
}
fn read_slow(
held: &Cached,
read_at: Timestamp,
read: impl Fn(Timestamp, &mut Vec<String>) -> Result<Vec<Device>>,
) -> Cached {
let mut warnings = Vec::new();
match read(read_at, &mut warnings) {
Ok(devices) => Cached {
devices,
warnings,
degraded: false,
},
Err(error) => Cached {
devices: held.devices.clone(),
warnings: vec![format!("{error}, keeping the last good reading")],
degraded: true,
},
}
}
fn read_fast(
read_at: Timestamp,
read: impl Fn(Timestamp, &mut Vec<String>) -> Vec<Device>,
cached: &Cached,
) -> Snapshot {
let mut warnings = cached.warnings.clone();
let devices = read(read_at, &mut warnings);
Snapshot {
degraded: cached.degraded,
..merge(devices, cached.devices.clone(), read_at, warnings)
}
}
const EARLY_FLOOR: Duration = Duration::from_secs(5);
fn slow_tier(
interval: Duration,
read: impl Fn(Timestamp, &mut Vec<String>) -> Result<Vec<Device>>,
clock: impl Fn() -> Timestamp,
refreshed: &Sender<Cached>,
wanted: &Receiver<()>,
) {
let mut held = Cached::default();
let mut early = false;
loop {
held = read_slow(&held, clock(), &read);
if refreshed.send(held.clone()).is_err() {
break;
}
if early {
thread::sleep(EARLY_FLOOR);
wanted.try_iter().for_each(drop);
}
match wanted.recv_timeout(interval) {
Ok(()) => early = true,
Err(RecvTimeoutError::Timeout) => early = false,
Err(RecvTimeoutError::Disconnected) => break,
}
}
}
fn fast_tier(
interval: Duration,
read: impl Fn(Timestamp, &mut Vec<String>) -> Vec<Device>,
clock: impl Fn() -> Timestamp,
snapshots: &Sender<Snapshot>,
cached: &Receiver<Cached>,
polling: Sender<()>,
nudges: &Receiver<()>,
) {
let mut latest = Cached::default();
loop {
latest = cached.try_iter().last().unwrap_or(latest);
if snapshots.send(read_fast(clock(), &read, &latest)).is_err() {
break;
}
if waited(nudges, interval) == Wake::Nudged {
let _ = polling.send(());
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Wake {
Tick,
Nudged,
}
fn waited(nudges: &Receiver<()>, interval: Duration) -> Wake {
match nudges.recv_timeout(interval) {
Ok(()) => {
nudges.try_iter().for_each(drop);
Wake::Nudged
}
Err(RecvTimeoutError::Timeout) => Wake::Tick,
Err(RecvTimeoutError::Disconnected) => {
thread::sleep(interval);
Wake::Tick
}
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::sync::atomic::{AtomicI64, Ordering};
use super::*;
use crate::address::Address;
use crate::device::{ChargeState, Levels, Source};
use crate::error::Error;
const READ_AT: Timestamp = Timestamp::from_unix(1_785_643_199);
const TRACKPAD: &str = "30-82-16-f2-24-90";
const KEYBOARD: &str = "de-df-38-f0-46-9b";
fn device(name: &str, address: &str, source: Source) -> Device {
Device {
address: Address::parse(address).expect("valid address"),
name: name.to_string(),
kind: None,
transport: None,
levels: Levels {
main: Some(85),
..Levels::default()
},
charge: ChargeState::Unknown,
source,
connected: true,
read_at: READ_AT,
}
}
fn trackpad() -> Device {
device("Magic Trackpad", TRACKPAD, Source::IoKit)
}
fn keyboard() -> Device {
device("MX Keys M Mac", KEYBOARD, Source::SystemProfiler)
}
fn frozen() -> impl Fn() -> Timestamp + Clone + Send + 'static {
|| READ_AT
}
fn unnudged() -> Receiver<()> {
mpsc::channel().1
}
fn counting_fast(
reads: Arc<AtomicI64>,
) -> impl Fn(Timestamp, &mut Vec<String>) -> Vec<Device> + Send + 'static {
move |_, _| {
vec![Device {
read_at: Timestamp::from_unix(reads.fetch_add(1, Ordering::SeqCst)),
..trackpad()
}]
}
}
fn counting_slow(
reads: Arc<AtomicI64>,
) -> impl Fn(Timestamp, &mut Vec<String>) -> Result<Vec<Device>> + Send + 'static {
move |_, _| {
reads.fetch_add(1, Ordering::SeqCst);
Ok(vec![keyboard()])
}
}
fn stamps(receiver: &Receiver<Snapshot>, count: usize) -> Vec<i64> {
receiver
.iter()
.take(count)
.map(|reading| reading.devices[0].read_at.unix())
.collect()
}
fn first(read: impl Fn(Timestamp, &mut Vec<String>) -> Result<Vec<Device>>) -> Cached {
read_slow(&Cached::default(), READ_AT, read)
}
fn failing(_: Timestamp, _: &mut Vec<String>) -> Result<Vec<Device>> {
Err(Error::Command("system_profiler exited with 1".to_string()))
}
#[test]
fn both_sources_merge_into_one_reading() {
let cached = first(|_, _| Ok(vec![keyboard()]));
let reading = read_fast(READ_AT, |_, _| vec![trackpad()], &cached);
assert_eq!(reading.devices.len(), 2);
assert!(reading.warnings.is_empty());
assert!(!reading.degraded);
}
#[test]
fn a_failed_system_profiler_degrades_the_reading_rather_than_failing_it() {
let cached = first(failing);
let reading = read_fast(READ_AT, |_, _| vec![trackpad()], &cached);
assert_eq!(reading.devices.len(), 1, "the fast source still answers");
assert!(reading.degraded);
assert_eq!(
reading.warnings,
["system_profiler exited with 1, keeping the last good reading"]
);
}
#[test]
fn a_failure_keeps_the_last_good_slow_devices_rather_than_dropping_them() {
let good = first(|_, _| Ok(vec![keyboard()]));
let degraded = read_slow(&good, READ_AT, failing);
let recovered = read_slow(°raded, READ_AT, |_, _| Ok(vec![keyboard()]));
let reading = read_fast(READ_AT, |_, _| vec![trackpad()], °raded);
assert_eq!(
reading.devices.len(),
2,
"the device only the slow source can see is still listed"
);
assert!(reading.degraded);
assert!(
!read_fast(READ_AT, |_, _| vec![trackpad()], &recovered).degraded,
"and the next good call clears it"
);
}
#[test]
fn a_cached_warning_travels_on_every_reading_it_applies_to() {
let cached = first(|_, warnings| {
warnings.push("skipped a malformed device".to_string());
Ok(Vec::new())
});
for _ in 0..3 {
let reading = read_fast(READ_AT, |_, _| vec![trackpad()], &cached);
assert_eq!(reading.warnings, ["skipped a malformed device"]);
}
}
#[test]
fn what_a_source_reports_this_tick_is_added_to_the_cached_warnings() {
let cached = first(|_, warnings| {
warnings.push("from the slow source".to_string());
Ok(Vec::new())
});
let reading = read_fast(
READ_AT,
|_, warnings| {
warnings.push("from the fast source".to_string());
Vec::new()
},
&cached,
);
assert_eq!(
reading.warnings,
["from the slow source", "from the fast source"]
);
}
#[test]
fn the_fast_tier_reads_immediately_and_then_on_every_interval() {
let reads = Arc::new(AtomicI64::new(0));
let receiver = poll_with(
Tiers {
fast: Duration::from_millis(1),
slow: Duration::from_secs(60),
..Tiers::default()
},
counting_fast(Arc::clone(&reads)),
|_, _| Ok(Vec::new()),
frozen(),
unnudged(),
);
assert_eq!(stamps(&receiver, 3), [0, 1, 2]);
}
#[test]
fn the_slow_tier_is_read_once_and_reused_across_fast_ticks() {
let slow_reads = Arc::new(AtomicI64::new(0));
let receiver = poll_with(
Tiers {
fast: Duration::from_millis(1),
slow: Duration::from_secs(60),
..Tiers::default()
},
|_, _| vec![trackpad()],
counting_slow(Arc::clone(&slow_reads)),
frozen(),
unnudged(),
);
let merged = receiver
.iter()
.take(500)
.position(|reading| reading.devices.len() == 2);
assert!(merged.is_some(), "the slow reading reaches a fast tick");
assert!(
receiver
.iter()
.take(5)
.all(|reading| reading.devices.len() == 2),
"and is reused on the ticks after it"
);
assert_eq!(slow_reads.load(Ordering::SeqCst), 1);
}
#[test]
fn a_nudge_reads_both_tiers_without_waiting_out_the_interval() {
let fast_reads = Arc::new(AtomicI64::new(0));
let slow_reads = Arc::new(AtomicI64::new(0));
let (nudge, nudges) = mpsc::channel();
let receiver = poll_with(
Tiers {
fast: Duration::from_secs(3_600),
slow: Duration::from_secs(3_600),
..Tiers::default()
},
counting_fast(Arc::clone(&fast_reads)),
counting_slow(Arc::clone(&slow_reads)),
frozen(),
nudges,
);
receiver.recv().expect("the first reading");
for _ in 0..3 {
nudge.send(()).expect("the poller is listening");
}
assert_eq!(
stamps(&receiver, 1),
[1],
"a second reading long before the hour is up"
);
assert!(
(0..500).any(|_| {
thread::sleep(Duration::from_millis(10));
slow_reads.load(Ordering::SeqCst) > 1
}),
"and the slow tier was asked to read again too"
);
}
#[test]
fn a_flapping_link_does_not_turn_into_a_stream_of_slow_reads() {
let slow_reads = Arc::new(AtomicI64::new(0));
let (nudge, nudges) = mpsc::channel();
let receiver = poll_with(
Tiers {
fast: Duration::from_secs(3_600),
slow: Duration::from_secs(3_600),
..Tiers::default()
},
|_, _| vec![trackpad()],
counting_slow(Arc::clone(&slow_reads)),
frozen(),
nudges,
);
receiver.recv().expect("the first reading");
for _ in 0..4 {
nudge.send(()).expect("the poller is listening");
thread::sleep(Duration::from_millis(50));
}
thread::sleep(Duration::from_millis(500));
assert_eq!(
slow_reads.load(Ordering::SeqCst),
2,
"one read on the first nudge, and the rest inside the floor"
);
}
#[test]
fn a_silent_nudge_source_leaves_the_tiers_on_their_intervals() {
let reads = Arc::new(AtomicI64::new(0));
let (_silent, nudges) = mpsc::channel();
let receiver = poll_with(
Tiers {
fast: Duration::from_millis(1),
slow: Duration::from_secs(60),
..Tiers::default()
},
counting_fast(Arc::clone(&reads)),
|_, _| Ok(Vec::new()),
frozen(),
nudges,
);
assert_eq!(stamps(&receiver, 3), [0, 1, 2]);
}
#[test]
fn a_hung_slow_source_never_delays_a_fast_reading() {
let (_blocked, never) = mpsc::channel::<()>();
let receiver = poll_with(
Tiers {
fast: Duration::from_millis(1),
slow: Duration::from_millis(1),
..Tiers::default()
},
counting_fast(Arc::new(AtomicI64::new(0))),
move |_, _| {
let _ = never.recv();
Ok(vec![keyboard()])
},
frozen(),
unnudged(),
);
assert_eq!(stamps(&receiver, 3), [0, 1, 2]);
}
#[test]
fn dropping_the_receiver_stops_both_tiers() {
let fast_reads = Arc::new(AtomicI64::new(0));
let slow_reads = Arc::new(AtomicI64::new(0));
let receiver = poll_with(
Tiers {
fast: Duration::from_millis(1),
slow: Duration::from_millis(1),
..Tiers::default()
},
counting_fast(Arc::clone(&fast_reads)),
counting_slow(Arc::clone(&slow_reads)),
frozen(),
unnudged(),
);
receiver.recv().expect("the first reading");
drop(receiver);
let stopped = (settled(&fast_reads), settled(&slow_reads));
thread::sleep(Duration::from_millis(50));
assert_eq!(
(
fast_reads.load(Ordering::SeqCst),
slow_reads.load(Ordering::SeqCst)
),
stopped,
"a stopped tier stays stopped"
);
}
fn settled(reads: &AtomicI64) -> i64 {
for _ in 0..500 {
let before = reads.load(Ordering::SeqCst);
thread::sleep(Duration::from_millis(10));
if reads.load(Ordering::SeqCst) == before {
return before;
}
}
panic!("the tier never stopped reading");
}
}