use std::collections::HashMap;
use std::fmt;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::mpsc;
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::{Duration, Instant};
use crate::capability::BackendKind;
use crate::error::{Error, Result};
use crate::ownership::ResourceId;
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum DnsEvent {
ResourceChanged {
resource: ResourceId,
},
ResourceRemoved {
resource: ResourceId,
},
}
impl DnsEvent {
pub(crate) fn resource(&self) -> &ResourceId {
match self {
DnsEvent::ResourceChanged { resource } => resource,
DnsEvent::ResourceRemoved { resource } => resource,
}
}
fn into_owned(self) -> (ResourceId, bool) {
match self {
DnsEvent::ResourceChanged { resource } => (resource, false),
DnsEvent::ResourceRemoved { resource } => (resource, true),
}
}
}
pub type WatchCallback = Arc<dyn Fn(&DnsEvent) + Send + Sync>;
pub struct WatchHandle {
flag: Arc<AtomicBool>,
cancel: Mutex<Option<Box<dyn FnOnce() + Send>>>,
}
impl WatchHandle {
#[allow(dead_code)]
pub(crate) fn new(flag: Arc<AtomicBool>, cancel: impl FnOnce() + Send + 'static) -> Self {
Self {
flag,
cancel: Mutex::new(Some(Box::new(cancel))),
}
}
pub(crate) fn is_active(&self) -> bool {
!self.flag.load(Ordering::Acquire)
}
pub fn stop(mut self) {
self.deactivate();
}
fn deactivate(&mut self) {
self.flag.store(true, Ordering::Release);
if let Some(cancel) = self
.cancel
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.take()
{
cancel();
}
}
}
impl Drop for WatchHandle {
fn drop(&mut self) {
self.deactivate();
}
}
impl fmt::Debug for WatchHandle {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("WatchHandle")
.field("active", &self.is_active())
.finish()
}
}
#[derive(Debug, Default)]
pub(crate) struct SuppressionRegistry {
entries: Mutex<HashMap<ResourceId, Instant>>,
}
impl SuppressionRegistry {
const WINDOW: Duration = Duration::from_millis(500);
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn suppress(&self, resource: &ResourceId) {
let mut entries = self
.entries
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
Self::prune(&mut entries);
entries.insert(resource.clone(), Instant::now());
}
pub(crate) fn is_suppressed(&self, resource: &ResourceId) -> bool {
let mut entries = self
.entries
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
Self::prune(&mut entries);
entries.contains_key(resource)
}
fn prune(entries: &mut HashMap<ResourceId, Instant>) {
let now = Instant::now();
entries.retain(|_, suppressed_at| now.duration_since(*suppressed_at) < Self::WINDOW);
}
}
pub(crate) fn spawn_coalescer(
kind: BackendKind,
callback: WatchCallback,
window: Duration,
) -> Result<Coalescer> {
let (tx, rx) = mpsc::channel::<DnsEvent>();
let worker = thread::Builder::new()
.name("osdns-watch-coalescer".to_string())
.spawn(move || {
let mut pending: HashMap<ResourceId, DnsEvent> = HashMap::new();
while let Ok(first) = rx.recv() {
let resource = first.resource().clone();
pending.insert(resource, first);
let deadline = Instant::now() + window;
loop {
match rx.recv_timeout(deadline.saturating_duration_since(Instant::now())) {
Ok(event) => {
let (resource, _) = event.clone().into_owned();
pending.insert(resource, event);
}
Err(mpsc::RecvTimeoutError::Timeout) => break,
Err(mpsc::RecvTimeoutError::Disconnected) => {
flush(&callback, &mut pending);
return;
}
}
}
flush(&callback, &mut pending);
}
})
.map_err(|e| Error::Platform {
backend: kind,
message: format!("cannot spawn coalescer thread: {e}"),
})?;
let sender: WatchCallback = Arc::new(move |event| {
let _ = tx.send(event.clone());
});
Ok(Coalescer {
sender: Some(sender),
worker: Some(worker),
})
}
pub(crate) struct Coalescer {
sender: Option<WatchCallback>,
worker: Option<thread::JoinHandle<()>>,
}
impl Coalescer {
pub(crate) fn callback(&self) -> WatchCallback {
self.sender.as_ref().expect("active coalescer").clone()
}
pub(crate) fn stop(mut self) {
self.sender.take();
if let Some(worker) = self.worker.take() {
let _ = worker.join();
}
}
}
impl Drop for Coalescer {
fn drop(&mut self) {
self.sender.take();
if let Some(worker) = self.worker.take() {
let _ = worker.join();
}
}
}
fn flush(callback: &WatchCallback, pending: &mut HashMap<ResourceId, DnsEvent>) {
for event in pending.values() {
callback(event);
}
pending.clear();
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn coalescer_delivers_the_last_observed_state() {
let delivered = Arc::new(Mutex::new(Vec::new()));
let sink = Arc::clone(&delivered);
let coalescer = spawn_coalescer(
BackendKind::Fake,
Arc::new(move |event| sink.lock().unwrap().push(event.clone())),
Duration::from_millis(5),
)
.unwrap();
let callback = coalescer.callback();
let resource = ResourceId::new("fake:watch-state").unwrap();
callback(&DnsEvent::ResourceRemoved {
resource: resource.clone(),
});
callback(&DnsEvent::ResourceChanged { resource });
drop(callback);
coalescer.stop();
assert!(matches!(
delivered.lock().unwrap().as_slice(),
[DnsEvent::ResourceChanged { .. }]
));
}
}