Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions crates/memtrack/src/ebpf/memtrack/maps.rs
Original file line number Diff line number Diff line change
Expand Up @@ -93,8 +93,8 @@ impl MemtrackBpf {
))
}

/// Callback that resumes every pressure-stopped process.
pub(super) fn on_ring_drained(&self) -> Box<dyn Fn() + Send> {
/// Callback that resumes pressure-stopped processes at the low watermark.
pub(super) fn on_ring_low_fill(&self) -> Box<dyn Fn() + Send> {
let stopped = self.stopped.clone();
Box::new(move || {
if let Err(error) = stopped.release_pressure() {
Expand Down
2 changes: 1 addition & 1 deletion crates/memtrack/src/ebpf/memtrack/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -282,7 +282,7 @@ impl MemtrackBpf {
resolve,
tx,
poll_interval_ms,
Some(self.on_ring_drained()),
Some(self.on_ring_low_fill()),
))
}

Expand Down
96 changes: 83 additions & 13 deletions crates/memtrack/src/ebpf/poller.rs
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,26 @@ fn consume_all(ringbuf: &RingBuffer, ring: *mut libbpf_sys::ring) {
}
}

/// Records read per regular-tick chunk. Bounds the work between fill checks.
const CONSUME_CHUNK_RECORDS: usize = BATCH_ITEMS;

/// Full chunks consumed per regular tick. Bounded so the tick ends even when
/// producers keep the ring full and control requests are re-checked between
/// ticks instead of being starved indefinitely; any backlog left over is
/// consumed on the next tick.
const MAX_CHUNKS_PER_TICK: usize = 64;

/// Consumes chunks until a short or failed chunk or [`MAX_CHUNKS_PER_TICK`],
/// running `on_full_chunk` after each full chunk.
fn consume_bounded_chunks(mut consume_chunk: impl FnMut() -> i32, mut on_full_chunk: impl FnMut()) {
for _ in 0..MAX_CHUNKS_PER_TICK {
if consume_chunk() != CONSUME_CHUNK_RECORDS as i32 {
return;
}
on_full_chunk();
}
}

fn poll_iteration<T>(
control: std::result::Result<Sender<()>, RecvTimeoutError>,
consume: impl FnOnce(),
Expand Down Expand Up @@ -72,8 +92,9 @@ fn poll_iteration<T>(
/// Polls a BPF ring buffer in a background thread, parsing raw entries with a
/// user-supplied closure and forwarding them to an mpsc channel in batches.
///
/// The poll thread runs until the poller is dropped, doing a final full
/// `consume()` on shutdown so no buffered entries are lost.
/// Regular ticks consume bounded chunks and check the low-fill release
/// watermark between chunks. Shutdown performs a final full `consume()` so no
/// buffered entries are lost.
pub struct RingBufferPoller {
ctl: Option<Sender<Sender<()>>>,
poll_thread: Option<JoinHandle<()>>,
Expand All @@ -85,7 +106,7 @@ impl RingBufferPoller {
parse: F,
tx: Sender<Vec<T>>,
poll_interval_ms: u64,
on_drained: Option<Box<dyn Fn() + Send>>,
on_low_fill: Option<Box<dyn Fn() + Send>>,
) -> Result<Self>
where
M: MapCore,
Expand Down Expand Up @@ -123,23 +144,33 @@ impl RingBufferPoller {
// SAFETY: the built `RingBuffer` holds exactly the one ring added above.
let ring =
unsafe { libbpf_sys::ring_buffer__ring(ringbuf.as_libbpf_object().as_ptr(), 0) };
// Resume below 1/4 fill; BPF stops at 3/4, which leaves hysteresis.
let release_if_low = || {
if let Some(on_low_fill) = &on_low_fill
&& unsafe { libbpf_sys::ring__avail_data_size(ring) }
< unsafe { libbpf_sys::ring__size(ring) } / 4
{
on_low_fill();
}
};
while poll_iteration(
ctl_rx.recv_timeout(Duration::from_millis(poll_interval_ms)),
|| consume_all(&ringbuf, ring),
|| {
let _ = ringbuf.poll(Duration::ZERO);
// A short or failed chunk ends the tick; the loop body below
// checks the fill after it.
consume_bounded_chunks(
|| ringbuf.consume_raw_n(CONSUME_CHUNK_RECORDS),
release_if_low,
);
},
&batch,
&tx,
) {
if let Some(on_drained) = &on_drained
&& unsafe { libbpf_sys::ring__avail_data_size(ring) } == 0
{
on_drained();
}
release_if_low();
}
if let Some(on_drained) = &on_drained {
on_drained();
if let Some(on_low_fill) = &on_low_fill {
on_low_fill();
}
});

Expand Down Expand Up @@ -192,7 +223,7 @@ impl ThreadedRingBufferPoller {
resolve: R,
tx: Sender<Vec<U>>,
poll_interval_ms: u64,
on_drained: Option<Box<dyn Fn() + Send>>,
on_low_fill: Option<Box<dyn Fn() + Send>>,
) -> Result<Self>
where
M: MapCore,
Expand All @@ -202,7 +233,7 @@ impl ThreadedRingBufferPoller {
R: Fn(T) -> U + Send + 'static,
{
let (parsed_tx, parsed_rx) = mpsc::channel::<Vec<T>>();
let ring = RingBufferPoller::new(rb_map, parse, parsed_tx, poll_interval_ms, on_drained)?;
let ring = RingBufferPoller::new(rb_map, parse, parsed_tx, poll_interval_ms, on_low_fill)?;
let resolver = std::thread::spawn(move || {
for batch in parsed_rx {
let resolved = batch.into_iter().map(&resolve).collect();
Expand Down Expand Up @@ -235,6 +266,45 @@ mod tests {
vec![42; BATCH_ITEMS - 1]
}

#[test]
fn bounded_chunks_end_at_cap_when_the_ring_stays_full() {
let reads = Cell::new(0);
let released = Cell::new(0);

consume_bounded_chunks(
|| {
reads.set(reads.get() + 1);
CONSUME_CHUNK_RECORDS as i32
},
|| released.set(released.get() + 1),
);

assert_eq!(reads.get(), MAX_CHUNKS_PER_TICK);
assert_eq!(released.get(), MAX_CHUNKS_PER_TICK);
}

#[test]
fn bounded_chunks_stop_on_short_chunk() {
let reads = Cell::new(0);
let released = Cell::new(0);

consume_bounded_chunks(
|| {
reads.set(reads.get() + 1);
if reads.get() == 3 {
// Short chunk: fewer records than requested.
0
} else {
CONSUME_CHUNK_RECORDS as i32
}
},
|| released.set(released.get() + 1),
);

assert_eq!(reads.get(), 3);
assert_eq!(released.get(), 2);
}

#[test]
fn timeout_flushes_partial_batch() {
let expected = partial_batch();
Expand Down
Loading