use choreo_proto::DaemonMessage;
use crossbeam_channel::Sender;
use std::collections::HashMap;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Clone)]
pub struct SubscriberSink {
pub tx: Sender<DaemonMessage>,
pub bytes_in_flight: Arc<AtomicUsize>,
}
impl SubscriberSink {
pub fn new(tx: Sender<DaemonMessage>) -> Self {
SubscriberSink {
tx,
bytes_in_flight: Arc::new(AtomicUsize::new(0)),
}
}
pub(crate) fn send_accounted(
&self,
msg: &DaemonMessage,
global: &AtomicUsize,
) -> Option<(usize, usize)> {
let size = msg.approx_wire_size();
let new_total = global.fetch_add(size, Ordering::Relaxed) + size;
let new_client = self.bytes_in_flight.fetch_add(size, Ordering::Relaxed) + size;
if self.tx.send(msg.clone()).is_ok() {
Some((new_client, new_total))
} else {
self.bytes_in_flight.fetch_sub(size, Ordering::Relaxed);
global.fetch_sub(size, Ordering::Relaxed);
None
}
}
pub fn enqueue(
&self,
msg: &DaemonMessage,
limits: &LagLimits,
global: &AtomicUsize,
) -> EnqueueOutcome {
let Some((new_client, new_total)) = self.send_accounted(msg, global) else {
return EnqueueOutcome::Disconnected;
};
if new_client > limits.per_client_cap {
EnqueueOutcome::ClientOverLag
} else if new_total > limits.global_budget {
EnqueueOutcome::GlobalOverBudget
} else {
EnqueueOutcome::Delivered
}
}
pub fn send_unchecked(&self, msg: &DaemonMessage, global: &AtomicUsize) -> bool {
self.send_accounted(msg, global).is_some()
}
}
pub(crate) fn fan_out_evicting(
subscribers: &mut HashMap<u64, SubscriberSink>,
msg: &DaemonMessage,
lag_limits: &LagLimits,
global: &AtomicUsize,
mut should_skip: impl FnMut(u64) -> bool,
) -> (Vec<u64>, bool) {
let mut evict_clients = Vec::new();
let mut evict_largest = false;
subscribers.retain(|client_id, sink| {
if should_skip(*client_id) {
return true;
}
match sink.enqueue(msg, lag_limits, global) {
EnqueueOutcome::Delivered => true,
EnqueueOutcome::Disconnected => false,
EnqueueOutcome::ClientOverLag => {
evict_clients.push(*client_id);
true
}
EnqueueOutcome::GlobalOverBudget => {
evict_largest = true;
true
}
}
});
(evict_clients, evict_largest)
}
#[cfg(test)]
pub(crate) fn test_sink() -> (SubscriberSink, crossbeam_channel::Receiver<DaemonMessage>) {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
(SubscriberSink::new(tx), rx)
}
pub enum EnqueueOutcome {
Delivered,
Disconnected,
ClientOverLag,
GlobalOverBudget,
}
#[derive(Debug, Clone, Copy)]
pub struct LagLimits {
pub per_client_cap: usize,
pub global_budget: usize,
}
impl Default for LagLimits {
fn default() -> Self {
LagLimits {
per_client_cap: 64 * 1024 * 1024,
global_budget: 512 * 1024 * 1024,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use choreo_proto::SessionEvent;
use choreo_proto::SessionStatus;
fn status_msg(session_id: u64) -> DaemonMessage {
DaemonMessage::Session {
session_id: Some(session_id),
event: SessionEvent::SessionStatusChanged {
status: SessionStatus::Inactive,
last_modified: 0,
},
}
}
fn tiny_limits() -> LagLimits {
LagLimits {
per_client_cap: 256,
global_budget: 256,
}
}
#[test]
fn enqueue_delivers_and_counts_bytes() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let sink = SubscriberSink::new(tx);
let global = AtomicUsize::new(0);
let limits = tiny_limits();
let msg = status_msg(1);
let size = msg.approx_wire_size();
let outcome = sink.enqueue(&msg, &limits, &global);
assert!(matches!(outcome, EnqueueOutcome::Delivered));
assert_eq!(rx.recv().unwrap(), msg, "message must be delivered");
assert_eq!(sink.bytes_in_flight.load(Ordering::Relaxed), size);
assert_eq!(global.load(Ordering::Relaxed), size);
sink.bytes_in_flight.fetch_sub(size, Ordering::Relaxed);
global.fetch_sub(size, Ordering::Relaxed);
assert_eq!(sink.bytes_in_flight.load(Ordering::Relaxed), 0);
assert_eq!(global.load(Ordering::Relaxed), 0);
}
#[test]
fn enqueue_over_per_client_cap_returns_client_over_lag_but_still_delivers() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let sink = SubscriberSink::new(tx);
let global = AtomicUsize::new(0);
let limits = LagLimits {
per_client_cap: 200,
global_budget: usize::MAX, };
let payload_msg = || DaemonMessage::Session {
session_id: Some(1),
event: SessionEvent::Failed {
request_id: 1,
error: "x".repeat(100),
},
};
let m1 = status_msg(1);
assert!(matches!(
sink.enqueue(&m1, &limits, &global),
EnqueueOutcome::Delivered
));
let m2 = payload_msg();
let outcome = sink.enqueue(&m2, &limits, &global);
assert!(
matches!(outcome, EnqueueOutcome::ClientOverLag),
"crossing the per-client cap must report ClientOverLag"
);
assert_eq!(rx.recv().unwrap(), m1);
assert_eq!(
rx.recv().unwrap(),
m2,
"the over-lag message is still enqueued"
);
}
#[test]
fn enqueue_over_global_budget_returns_global_over_budget_but_still_delivers() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let sink = SubscriberSink::new(tx);
let global = AtomicUsize::new(0);
let limits = LagLimits {
per_client_cap: usize::MAX,
global_budget: 200,
};
let m1 = status_msg(1);
assert!(matches!(
sink.enqueue(&m1, &limits, &global),
EnqueueOutcome::Delivered
));
let m2 = DaemonMessage::Session {
session_id: Some(1),
event: SessionEvent::Failed {
request_id: 2,
error: "y".repeat(200),
},
};
let outcome = sink.enqueue(&m2, &limits, &global);
assert!(
matches!(outcome, EnqueueOutcome::GlobalOverBudget),
"crossing the global budget must report GlobalOverBudget"
);
assert_eq!(rx.recv().unwrap(), m1);
assert_eq!(
rx.recv().unwrap(),
m2,
"the over-budget message is still enqueued"
);
}
#[test]
fn enqueue_returns_disconnected_when_receiver_gone() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let sink = SubscriberSink::new(tx);
let global = AtomicUsize::new(0);
let limits = tiny_limits();
drop(rx);
let msg = status_msg(1);
let outcome = sink.enqueue(&msg, &limits, &global);
assert!(matches!(outcome, EnqueueOutcome::Disconnected));
assert_eq!(sink.bytes_in_flight.load(Ordering::Relaxed), 0);
assert_eq!(global.load(Ordering::Relaxed), 0);
}
#[test]
fn send_unchecked_self_corrects_on_dead_receiver() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let sink = SubscriberSink::new(tx);
let global = AtomicUsize::new(0);
drop(rx);
let msg = status_msg(1);
let size = msg.approx_wire_size();
assert!(
!sink.send_unchecked(&msg, &global),
"dead receiver must report false"
);
assert_eq!(sink.bytes_in_flight.load(Ordering::Relaxed), 0);
assert_eq!(global.load(Ordering::Relaxed), 0);
let (tx2, rx2) = crossbeam_channel::unbounded::<DaemonMessage>();
let sink2 = SubscriberSink::new(tx2);
let global2 = AtomicUsize::new(0);
assert!(
sink2.send_unchecked(&msg, &global2),
"live receiver must report true"
);
assert_eq!(rx2.recv().unwrap(), msg, "message must be delivered");
assert_eq!(sink2.bytes_in_flight.load(Ordering::Relaxed), size);
assert_eq!(global2.load(Ordering::Relaxed), size);
}
#[test]
fn counters_increment_and_decrement_with_approx_wire_size() {
let (tx, rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let sink = SubscriberSink::new(tx);
let global = AtomicUsize::new(0);
let limits = LagLimits {
per_client_cap: usize::MAX,
global_budget: usize::MAX,
};
let m1 = DaemonMessage::Session {
session_id: Some(1),
event: SessionEvent::Failed {
request_id: 1,
error: "a".repeat(100),
},
};
let m2 = DaemonMessage::Session {
session_id: Some(2),
event: SessionEvent::Failed {
request_id: 2,
error: "b".repeat(50),
},
};
let s1 = m1.approx_wire_size();
let s2 = m2.approx_wire_size();
assert!(matches!(
sink.enqueue(&m1, &limits, &global),
EnqueueOutcome::Delivered
));
assert!(matches!(
sink.enqueue(&m2, &limits, &global),
EnqueueOutcome::Delivered
));
assert_eq!(sink.bytes_in_flight.load(Ordering::Relaxed), s1 + s2);
assert_eq!(global.load(Ordering::Relaxed), s1 + s2);
assert_eq!(rx.recv().unwrap(), m1);
assert_eq!(rx.recv().unwrap(), m2);
sink.bytes_in_flight.fetch_sub(s1, Ordering::Relaxed);
sink.bytes_in_flight.fetch_sub(s2, Ordering::Relaxed);
global.fetch_sub(s1, Ordering::Relaxed);
global.fetch_sub(s2, Ordering::Relaxed);
assert_eq!(sink.bytes_in_flight.load(Ordering::Relaxed), 0);
assert_eq!(global.load(Ordering::Relaxed), 0);
}
#[test]
fn default_limits_are_64_mib_per_client_and_512_mib_global() {
let limits = LagLimits::default();
assert_eq!(limits.per_client_cap, 64 * 1024 * 1024);
assert_eq!(limits.global_budget, 512 * 1024 * 1024);
}
#[test]
fn per_client_overlag_takes_precedence_over_global_over_budget() {
let (tx, _rx) = crossbeam_channel::unbounded::<DaemonMessage>();
let sink = SubscriberSink::new(tx);
let global = AtomicUsize::new(0);
let limits = LagLimits {
per_client_cap: 8,
global_budget: 8,
};
let outcome = sink.enqueue(&status_msg(1), &limits, &global);
assert!(matches!(outcome, EnqueueOutcome::ClientOverLag));
}
}