use std::collections::HashMap;
use std::sync::Mutex;
#[cfg(test)]
use std::sync::atomic::{AtomicBool, Ordering};
use serde::Serialize;
use tau_cli_term::RendererDeliveryId;
use tau_delivery_memory::DecodedMemoryEstimate;
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(super) enum DeliveryMemoryCut {
DecodeCurrent,
ColdStaging,
RendererFifo,
Scheduler,
Handler,
}
impl DeliveryMemoryCut {
const COUNT: usize = 5;
const ALL: [Self; Self::COUNT] = [
Self::DecodeCurrent,
Self::ColdStaging,
Self::RendererFifo,
Self::Scheduler,
Self::Handler,
];
const fn label(self) -> &'static str {
match self {
Self::DecodeCurrent => "decode_current",
Self::ColdStaging => "cold_staging",
Self::RendererFifo => "renderer_fifo",
Self::Scheduler => "scheduler",
Self::Handler => "handler",
}
}
const fn index(self) -> usize {
self as usize
}
}
#[derive(Clone, Copy)]
struct ActiveEstimate {
cut: DeliveryMemoryCut,
estimate: DecodedMemoryEstimate,
}
pub(super) struct DeliveryMemoryTracker {
state: Mutex<Option<Box<TrackerState>>>,
#[cfg(test)]
force_enabled: AtomicBool,
}
struct TrackerState {
active: HashMap<RendererDeliveryId, ActiveEstimate>,
high_water: [DecodedMemoryEstimate; DeliveryMemoryCut::COUNT],
high_water_items: [u64; DeliveryMemoryCut::COUNT],
}
impl DeliveryMemoryTracker {
pub(super) const fn new() -> Self {
Self {
state: Mutex::new(None),
#[cfg(test)]
force_enabled: AtomicBool::new(false),
}
}
pub(super) fn observe_decode(
&self,
delivery_id: RendererDeliveryId,
message: &impl Serialize,
encoded_bytes: tau_proto::ProtocolMessageBytes,
) {
if !self.enabled() {
return;
}
let Some(estimate) = DecodedMemoryEstimate::from_serializable(message, encoded_bytes)
else {
return;
};
let mut state = self.state.lock().expect("delivery-memory mutex poisoned");
let state = state.get_or_insert_with(|| {
Box::new(TrackerState {
active: HashMap::new(),
high_water: [DecodedMemoryEstimate::default(); DeliveryMemoryCut::COUNT],
high_water_items: [0; DeliveryMemoryCut::COUNT],
})
});
state.active.insert(
delivery_id,
ActiveEstimate {
cut: DeliveryMemoryCut::DecodeCurrent,
estimate,
},
);
state.emit_snapshot();
}
pub(super) fn transition(&self, delivery_id: RendererDeliveryId, cut: DeliveryMemoryCut) {
if !self.enabled() {
return;
}
let mut state = self.state.lock().expect("delivery-memory mutex poisoned");
let Some(state) = state.as_mut() else {
return;
};
let Some(estimate) = state.active.get_mut(&delivery_id) else {
return;
};
if cut == DeliveryMemoryCut::RendererFifo
&& matches!(
estimate.cut,
DeliveryMemoryCut::Scheduler | DeliveryMemoryCut::Handler
)
{
return;
}
estimate.cut = cut;
state.emit_snapshot();
}
pub(super) fn release(&self, delivery_id: RendererDeliveryId) {
if !self.enabled() {
return;
}
let mut state = self.state.lock().expect("delivery-memory mutex poisoned");
let Some(current) = state.as_mut() else {
return;
};
current.active.remove(&delivery_id);
current.emit_snapshot();
}
fn enabled(&self) -> bool {
#[cfg(test)]
if self.force_enabled.load(Ordering::Relaxed) {
return true;
}
tracing::enabled!(target: "tau_cli::delivery_memory", tracing::Level::TRACE)
}
#[cfg(test)]
pub(super) fn force_enable_for_test(&self) {
self.force_enabled.store(true, Ordering::Relaxed);
}
#[cfg(test)]
pub(super) fn cut_for_test(
&self,
delivery_id: RendererDeliveryId,
) -> Option<DeliveryMemoryCut> {
self.state
.lock()
.expect("delivery-memory mutex poisoned")
.as_ref()
.and_then(|state| state.active.get(&delivery_id))
.map(|active| active.cut)
}
#[cfg(test)]
pub(super) fn active_len_for_test(&self) -> usize {
self.state
.lock()
.expect("delivery-memory mutex poisoned")
.as_ref()
.map_or(0, |state| state.active.len())
}
}
impl TrackerState {
fn emit_snapshot(&mut self) {
for cut in DeliveryMemoryCut::ALL {
let (items, estimate) = self.active.values().filter(|item| item.cut == cut).fold(
(0_u64, DecodedMemoryEstimate::default()),
|(items, total), item| {
(items.saturating_add(1), total.saturating_add(item.estimate))
},
);
let index = cut.index();
self.high_water[index].encoded_bytes = self.high_water[index]
.encoded_bytes
.max(estimate.encoded_bytes);
self.high_water[index].logical_payload_bytes = self.high_water[index]
.logical_payload_bytes
.max(estimate.logical_payload_bytes);
self.high_water[index].requested_capacity_estimate = self.high_water[index]
.requested_capacity_estimate
.max(estimate.requested_capacity_estimate);
self.high_water[index].container_count = self.high_water[index]
.container_count
.max(estimate.container_count);
self.high_water_items[index] = self.high_water_items[index].max(items);
let high = self.high_water[index];
tracing::trace!(
target: "tau_cli::delivery_memory",
process = "cli",
cut = cut.label(),
items,
owners = u64::from(items != 0),
encoded_bytes = estimate.encoded_bytes.get(),
decoded_logical_bytes_estimate = estimate.logical_payload_bytes.get(),
decoded_requested_capacity_estimate = estimate.requested_capacity_estimate.get(),
decoded_containers = estimate.container_count,
expansion_milli = estimate.expansion_milli(),
shared_allocations = items,
shared_fanout = 0_u64,
high_water_items = self.high_water_items[index],
high_water_encoded_bytes = high.encoded_bytes.get(),
high_water_decoded_logical_bytes_estimate = high.logical_payload_bytes.get(),
high_water_decoded_requested_capacity_estimate = high.requested_capacity_estimate.get(),
kernel_bytes_observable = false,
retained_projection_bytes_observable = false,
"decoded delivery memory ownership"
);
}
}
}
#[cfg(test)]
mod tests;