use std::sync::OnceLock;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use super::cgroup;
use super::usage::{UsageReader, UsageSource};
static HEAP_SOURCE: OnceLock<fn() -> usize> = OnceLock::new();
#[must_use]
pub fn set_heap_source(source: fn() -> usize) -> bool {
HEAP_SOURCE.set(source).is_ok()
}
#[inline]
fn heap_bytes() -> Option<u64> {
HEAP_SOURCE.get().map(|f| f() as u64)
}
fn env_parsed<T: std::str::FromStr>(prefix: &str, suffix: &str) -> Option<T> {
std::env::var(format!("{prefix}_{suffix}"))
.ok()
.and_then(|v| v.parse().ok())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MemoryPressure {
Low,
Medium,
High,
}
#[derive(Debug, Clone, serde::Deserialize, serde::Serialize)]
#[cfg_attr(feature = "config-schema", derive(schemars::JsonSchema))]
pub struct MemoryGuardConfig {
#[serde(default)]
pub limit_bytes: u64,
#[serde(default = "default_pressure_threshold")]
pub pressure_threshold: f64,
#[serde(default = "default_cgroup_headroom")]
pub cgroup_headroom: f64,
}
fn default_pressure_threshold() -> f64 {
DEFAULT_PRESSURE_THRESHOLD
}
fn default_cgroup_headroom() -> f64 {
DEFAULT_CGROUP_HEADROOM
}
fn check_fraction(v: f64, name: &str) -> Result<(), String> {
if !v.is_finite() || v <= 0.0 || v > 1.0 {
return Err(format!(
"memory.{name} must be a finite fraction in (0.0, 1.0], got {v}"
));
}
Ok(())
}
fn sane_fraction(v: f64, default: f64, name: &str) -> f64 {
if check_fraction(v, name).is_err() {
tracing::error!(
value = v,
"invalid memory.{name} (need finite fraction in (0,1]); using default {default}"
);
default
} else {
v
}
}
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
fn effective_auto_limit(detected: u64, headroom: f64, high: Option<u64>) -> u64 {
let headroom_limit = (detected as f64 * headroom) as u64;
match high {
Some(h) => headroom_limit.min(h),
None => headroom_limit,
}
}
const DEFAULT_CGROUP_HEADROOM: f64 = 0.85;
const DEFAULT_PRESSURE_THRESHOLD: f64 = 0.80;
impl Default for MemoryGuardConfig {
fn default() -> Self {
Self {
limit_bytes: 0, pressure_threshold: DEFAULT_PRESSURE_THRESHOLD,
cgroup_headroom: DEFAULT_CGROUP_HEADROOM,
}
}
}
impl MemoryGuardConfig {
#[must_use]
pub fn from_cascade() -> Self {
#[cfg(feature = "config")]
{
if let Some(cfg) = crate::config::try_get()
&& let Ok(memory) = cfg.unmarshal_key_registered::<Self>("memory")
{
return memory;
}
}
Self::default()
}
#[must_use]
#[cfg(feature = "config")]
pub fn from_env(prefix: &str) -> Self {
use crate::config::flat_env::flat_env_parsed;
let mut config = Self::default();
if let Some(v) = flat_env_parsed::<u64>(prefix, "MEMORY_LIMIT_BYTES") {
config.limit_bytes = v;
}
if let Some(v) = flat_env_parsed::<f64>(prefix, "MEMORY_PRESSURE_THRESHOLD") {
config.pressure_threshold = v;
}
if let Some(v) = flat_env_parsed::<f64>(prefix, "MEMORY_CGROUP_HEADROOM") {
config.cgroup_headroom = v;
}
config
}
#[must_use]
pub fn from_env_raw(prefix: &str) -> Self {
let mut config = Self::default();
if let Some(v) = env_parsed::<u64>(prefix, "MEMORY_LIMIT_BYTES") {
config.limit_bytes = v;
}
if let Some(v) = env_parsed::<f64>(prefix, "MEMORY_PRESSURE_THRESHOLD") {
config.pressure_threshold = v;
}
if let Some(v) = env_parsed::<f64>(prefix, "MEMORY_CGROUP_HEADROOM") {
config.cgroup_headroom = v;
}
config
}
pub fn validate(&self) -> Result<(), String> {
check_fraction(self.pressure_threshold, "pressure_threshold")?;
check_fraction(self.cgroup_headroom, "cgroup_headroom")?;
Ok(())
}
}
pub struct MemoryGuard {
reserved_bytes: AtomicU64,
usage: UsageReader,
limit_bytes: u64,
pressure_threshold: f64,
under_pressure: AtomicBool,
}
impl MemoryGuard {
#[must_use]
pub fn new(config: MemoryGuardConfig) -> Self {
Self::with_usage_source(config, UsageSource::detect())
}
#[must_use]
#[allow(clippy::cast_possible_truncation, clippy::cast_sign_loss)]
pub fn with_usage_source(config: MemoryGuardConfig, source: UsageSource) -> Self {
let pressure_threshold = sane_fraction(
config.pressure_threshold,
DEFAULT_PRESSURE_THRESHOLD,
"pressure_threshold",
);
let cgroup_headroom = sane_fraction(
config.cgroup_headroom,
DEFAULT_CGROUP_HEADROOM,
"cgroup_headroom",
);
let raw_limit = if config.limit_bytes > 0 {
config.limit_bytes
} else {
effective_auto_limit(
cgroup::detect_memory_limit(),
cgroup_headroom,
cgroup::detect_memory_high(),
)
};
let limit_bytes = raw_limit.max(1);
let usage = UsageReader::new(source);
let usage_source = if HEAP_SOURCE.get().is_some() {
"explicit"
} else {
usage.source().name()
};
let usage_bytes = heap_bytes().or_else(|| usage.read());
tracing::info!(
limit_bytes,
pressure_threshold,
usage_source,
?usage_bytes,
"memory guard initialised"
);
if usage_bytes.is_none() {
tracing::warn!(
"no cgroup or procfs memory accounting readable: the guard sees \
only bytes callers reserve and release, not the process's usage"
);
}
Self {
reserved_bytes: AtomicU64::new(0),
usage,
limit_bytes,
pressure_threshold,
under_pressure: AtomicBool::new(false),
}
}
#[inline]
pub fn try_reserve(&self, bytes: u64) -> bool {
if let Some(heap) = heap_bytes() {
return heap.saturating_add(bytes) <= self.limit_bytes;
}
if let Some(estimate) = self.usage.estimate() {
if estimate.saturating_add(bytes) > self.limit_bytes {
return false;
}
self.usage.admit(bytes);
return true;
}
let current = self.reserved_bytes.fetch_add(bytes, Ordering::Relaxed) + bytes;
if current > self.limit_bytes {
self.reserved_bytes.fetch_sub(bytes, Ordering::Relaxed);
self.under_pressure.store(true, Ordering::Relaxed);
return false;
}
self.update_pressure(current);
true
}
#[inline]
pub fn add_bytes(&self, bytes: u64) {
self.usage.admit(bytes);
let new_total = self.reserved_bytes.fetch_add(bytes, Ordering::Relaxed) + bytes;
self.update_pressure(new_total);
}
#[inline]
pub fn release(&self, bytes: u64) {
self.usage.forget(bytes);
let prev = self
.reserved_bytes
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |current| {
Some(current.saturating_sub(bytes))
})
.unwrap_or_else(|v| v);
self.update_pressure(prev.saturating_sub(bytes));
}
#[inline]
pub fn under_pressure(&self) -> bool {
if self.usage_bytes().is_some() {
return self.pressure_ratio() >= self.pressure_threshold;
}
self.under_pressure.load(Ordering::Relaxed)
}
#[inline]
pub fn pressure(&self) -> MemoryPressure {
let ratio = self.pressure_ratio();
if ratio >= self.pressure_threshold {
MemoryPressure::High
} else if ratio >= 0.5 {
MemoryPressure::Medium
} else {
MemoryPressure::Low
}
}
#[inline]
pub fn pressure_ratio(&self) -> f64 {
self.current_bytes() as f64 / self.limit_bytes as f64
}
#[inline]
pub fn current_bytes(&self) -> u64 {
self.usage_bytes()
.unwrap_or_else(|| self.reserved_bytes.load(Ordering::Relaxed))
}
#[inline]
pub fn reserved_bytes(&self) -> u64 {
self.reserved_bytes.load(Ordering::Relaxed)
}
#[must_use]
pub fn usage_source(&self) -> &'static str {
if HEAP_SOURCE.get().is_some() {
"explicit"
} else {
self.usage.source().name()
}
}
#[inline]
fn usage_bytes(&self) -> Option<u64> {
heap_bytes().or_else(|| self.usage.estimate())
}
#[inline]
pub fn limit_bytes(&self) -> u64 {
self.limit_bytes
}
#[inline]
fn update_pressure(&self, current: u64) {
let ratio = current as f64 / self.limit_bytes as f64;
self.under_pressure
.store(ratio >= self.pressure_threshold, Ordering::Relaxed);
}
}
#[cfg(test)]
mod tests {
use super::*;
fn reservation_guard(config: MemoryGuardConfig) -> MemoryGuard {
MemoryGuard::with_usage_source(config, UsageSource::Reservations)
}
fn cgroup_fixture(files: &[(&str, &str)]) -> tempfile::TempDir {
let dir = tempfile::tempdir().expect("tempdir");
for (name, contents) in files {
std::fs::write(dir.path().join(name), contents).expect("write fixture");
}
dir
}
#[test]
fn test_memory_guard_default() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 1_000_000, ..Default::default()
});
assert_eq!(guard.limit_bytes(), 1_000_000);
assert_eq!(guard.current_bytes(), 0);
assert!(!guard.under_pressure());
assert_eq!(guard.pressure(), MemoryPressure::Low);
}
#[test]
fn test_try_reserve_within_limit() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 1000,
..Default::default()
});
assert!(guard.try_reserve(500));
assert_eq!(guard.reserved_bytes(), 500);
}
#[test]
fn test_try_reserve_over_limit() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 1000,
..Default::default()
});
assert!(guard.try_reserve(500));
assert!(!guard.try_reserve(600)); assert_eq!(guard.reserved_bytes(), 500); assert!(guard.under_pressure());
}
#[test]
fn test_release_reduces_pressure() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 1000,
pressure_threshold: 0.8,
..Default::default()
});
guard.add_bytes(900); assert!(guard.under_pressure());
assert_eq!(guard.pressure(), MemoryPressure::High);
guard.release(500); assert!(!guard.under_pressure());
assert_eq!(guard.pressure(), MemoryPressure::Low);
}
#[test]
fn test_pressure_levels() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 1000,
pressure_threshold: 0.8,
..Default::default()
});
guard.add_bytes(400);
assert_eq!(guard.pressure(), MemoryPressure::Low);
guard.add_bytes(200); assert_eq!(guard.pressure(), MemoryPressure::Medium);
guard.add_bytes(300); assert_eq!(guard.pressure(), MemoryPressure::High);
}
#[test]
fn test_pressure_ratio() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 1000,
..Default::default()
});
guard.add_bytes(250);
let ratio = guard.pressure_ratio();
assert!((ratio - 0.25).abs() < 0.001);
}
#[test]
fn test_release_saturating() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 1000,
..Default::default()
});
guard.add_bytes(100);
guard.release(200); assert_eq!(
guard.reserved_bytes(),
0,
"over-release must saturate to 0, not wrap"
);
assert!(!guard.under_pressure());
assert_eq!(guard.pressure(), MemoryPressure::Low);
assert!(guard.try_reserve(500));
assert_eq!(guard.reserved_bytes(), 500);
}
#[test]
fn test_concurrent_reserve_release() {
use std::sync::Arc;
use std::thread;
let guard = Arc::new(reservation_guard(MemoryGuardConfig {
limit_bytes: 100_000,
pressure_threshold: 0.8,
..Default::default()
}));
let mut handles = vec![];
for _ in 0..10 {
let g = Arc::clone(&guard);
handles.push(thread::spawn(move || {
for _ in 0..100 {
g.add_bytes(100);
g.release(100);
}
}));
}
for h in handles {
h.join().unwrap();
}
assert!(
guard.reserved_bytes() < 1000,
"leaked bytes: {}",
guard.reserved_bytes()
);
}
#[test]
fn test_try_reserve_rollback_is_atomic() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 100,
..Default::default()
});
assert!(guard.try_reserve(90));
assert!(!guard.try_reserve(20)); assert_eq!(guard.reserved_bytes(), 90); assert!(guard.try_reserve(10)); assert_eq!(guard.reserved_bytes(), 100);
}
#[test]
fn pressure_ratio_follows_the_cgroup_current_file() {
let limit = 536_870_912u64; let dir = cgroup_fixture(&[("memory.current", "0\n"), ("memory.max", "536870912\n")]);
let guard = MemoryGuard::with_usage_source(
MemoryGuardConfig {
limit_bytes: limit,
pressure_threshold: 0.8,
..Default::default()
},
UsageSource::CgroupV2(dir.path().to_path_buf()),
);
std::fs::write(dir.path().join("memory.current"), "198299648\n").expect("write");
std::thread::sleep(std::time::Duration::from_millis(60)); let ratio = guard.pressure_ratio();
assert!(
(ratio - 0.369).abs() < 0.01,
"189 MiB of 512 MiB is 0.369, got {ratio}"
);
assert_eq!(guard.current_bytes(), 198_299_648);
assert_eq!(guard.pressure(), MemoryPressure::Low);
assert!(!guard.under_pressure());
std::fs::write(dir.path().join("memory.current"), "471859200\n").expect("write");
std::thread::sleep(std::time::Duration::from_millis(60));
assert!(
guard.under_pressure(),
"88% of the cgroup limit is over the 80% threshold"
);
assert_eq!(guard.pressure(), MemoryPressure::High);
assert!(
!guard.try_reserve(64 * 1024 * 1024),
"admission is projected against the cgroup, so 450 + 64 MiB is refused"
);
}
fn ledger_guard(dir: &tempfile::TempDir) -> MemoryGuard {
MemoryGuard::with_usage_source(
MemoryGuardConfig {
limit_bytes: 10 * 1024 * 1024,
..Default::default()
},
UsageSource::CgroupV2(dir.path().to_path_buf()),
)
}
#[test]
fn a_burst_inside_one_cache_window_is_refused() {
let dir = cgroup_fixture(&[("memory.current", "0\n")]);
let guard = ledger_guard(&dir);
let six_mib = 6 * 1024 * 1024;
assert!(guard.try_reserve(six_mib), "6 MiB of a 10 MiB limit fits");
assert!(
!guard.try_reserve(six_mib),
"a second 6 MiB is 12 MiB of a 10 MiB limit, against an unchanged file"
);
}
#[test]
fn release_discharges_the_ledger() {
let dir = cgroup_fixture(&[("memory.current", "0\n")]);
let guard = ledger_guard(&dir);
let six_mib = 6 * 1024 * 1024;
assert!(guard.try_reserve(six_mib));
guard.release(six_mib);
assert!(
guard.try_reserve(six_mib),
"the released bytes no longer count against the next admission"
);
}
#[test]
fn a_fresh_sample_clears_the_ledger() {
let dir = cgroup_fixture(&[("memory.current", "0\n")]);
let guard = ledger_guard(&dir);
let six_mib = 6 * 1024 * 1024;
assert!(guard.try_reserve(six_mib));
std::thread::sleep(std::time::Duration::from_millis(60)); assert!(
guard.try_reserve(six_mib),
"the new reading accounts for those bytes, so the ledger restarts"
);
}
#[test]
fn the_ledger_adds_to_what_the_kernel_already_charges() {
let dir = cgroup_fixture(&[("memory.current", "4194304\n")]); let guard = ledger_guard(&dir);
assert_eq!(guard.current_bytes(), 4 * 1024 * 1024);
assert!(guard.try_reserve(2 * 1024 * 1024));
assert_eq!(
guard.current_bytes(),
6 * 1024 * 1024,
"4 MiB charged plus 2 MiB admitted against that reading"
);
}
#[test]
fn reserved_bytes_and_current_bytes_are_different_numbers() {
let dir = cgroup_fixture(&[("memory.current", "104857600\n")]);
let guard = MemoryGuard::with_usage_source(
MemoryGuardConfig {
limit_bytes: 536_870_912,
..Default::default()
},
UsageSource::CgroupV2(dir.path().to_path_buf()),
);
guard.add_bytes(4096);
assert_eq!(
guard.reserved_bytes(),
4096,
"the caller's outstanding lease"
);
assert_eq!(
guard.current_bytes(),
104_861_696,
"what the kernel charges, plus the lease taken against that reading"
);
}
#[test]
fn usage_source_is_named_for_the_init_log() {
let dir = cgroup_fixture(&[("memory.current", "1024\n")]);
let guard = MemoryGuard::with_usage_source(
MemoryGuardConfig {
limit_bytes: 1000,
..Default::default()
},
UsageSource::CgroupV2(dir.path().to_path_buf()),
);
assert_eq!(guard.usage_source(), "cgroup-v2");
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 1000,
..Default::default()
});
assert_eq!(guard.usage_source(), "reservations");
}
static TEST_HEAP: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
fn test_heap_source() -> usize {
TEST_HEAP.load(Ordering::Relaxed)
}
#[test]
fn heap_source_overrides_read_path_and_admission() {
assert!(set_heap_source(test_heap_source), "first set wins");
assert!(
!set_heap_source(test_heap_source),
"second set is a no-op (first-wins)"
);
let guard = MemoryGuard::new(MemoryGuardConfig {
limit_bytes: 1_000,
pressure_threshold: 0.8,
..Default::default()
});
TEST_HEAP.store(250, Ordering::Relaxed);
assert_eq!(guard.current_bytes(), 250);
assert!((guard.pressure_ratio() - 0.25).abs() < 0.001);
assert!(!guard.under_pressure());
TEST_HEAP.store(850, Ordering::Relaxed);
assert!(
guard.under_pressure(),
"85% live heap is over the 80% threshold"
);
assert_eq!(guard.pressure(), MemoryPressure::High);
TEST_HEAP.store(900, Ordering::Relaxed);
assert!(guard.try_reserve(100), "900 + 100 == limit, admitted");
assert!(!guard.try_reserve(200), "900 + 200 > limit, rejected");
assert_eq!(guard.current_bytes(), 900, "still the heap source");
assert_eq!(
guard.reserved_bytes(),
0,
"counter untouched by try_reserve"
);
assert_eq!(guard.usage_source(), "explicit");
}
#[test]
fn effective_auto_limit_caps_at_memory_high() {
assert_eq!(effective_auto_limit(1000, 0.85, Some(600)), 600);
assert_eq!(effective_auto_limit(1000, 0.85, Some(900)), 850);
assert_eq!(effective_auto_limit(1000, 0.85, None), 850);
}
#[test]
fn test_config_defaults() {
let config = MemoryGuardConfig::default();
assert_eq!(config.limit_bytes, 0);
assert!((config.pressure_threshold - 0.80).abs() < 0.001);
assert!((config.cgroup_headroom - 0.85).abs() < 0.001);
}
#[test]
fn test_from_env_raw_defaults_when_unset() {
let config = MemoryGuardConfig::from_env_raw("TEST_MG_UNSET");
assert_eq!(config.limit_bytes, 0);
assert!((config.pressure_threshold - 0.80).abs() < 0.001);
assert!((config.cgroup_headroom - 0.85).abs() < 0.001);
}
#[test]
fn test_env_parsed_helper() {
assert!(env_parsed::<u64>("NONEXISTENT_PREFIX_XYZ", "FOO").is_none());
assert!(env_parsed::<f64>("NONEXISTENT_PREFIX_XYZ", "BAR").is_none());
}
#[test]
fn test_guard_with_explicit_config_overrides() {
let config = MemoryGuardConfig {
limit_bytes: 2_147_483_648,
pressure_threshold: 0.75,
cgroup_headroom: 0.90,
};
let guard = MemoryGuard::new(config);
assert_eq!(guard.limit_bytes(), 2_147_483_648);
}
#[test]
fn test_guard_with_custom_headroom() {
let config = MemoryGuardConfig {
limit_bytes: 0, pressure_threshold: 0.80,
cgroup_headroom: 0.85,
};
let guard = MemoryGuard::new(config);
assert!(guard.limit_bytes() > 0);
}
#[test]
fn test_validate_accepts_defaults_and_rejects_bad_fractions() {
assert!(MemoryGuardConfig::default().validate().is_ok());
for bad in [0.0, -0.1, 1.5, f64::NAN, f64::INFINITY] {
let cfg = MemoryGuardConfig {
pressure_threshold: bad,
..Default::default()
};
assert!(
cfg.validate().is_err(),
"pressure_threshold={bad} must be rejected"
);
let cfg = MemoryGuardConfig {
cgroup_headroom: bad,
..Default::default()
};
assert!(
cfg.validate().is_err(),
"cgroup_headroom={bad} must be rejected"
);
}
}
#[test]
fn test_new_clamps_invalid_config_no_divide_by_zero() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 0,
pressure_threshold: 0.0,
cgroup_headroom: 0.0,
});
assert!(guard.limit_bytes() >= 1, "limit floored at >=1");
guard.add_bytes(10);
assert!(
guard.pressure_ratio().is_finite(),
"pressure ratio must be finite, not div-by-zero"
);
}
#[test]
fn test_new_with_nan_threshold_is_finite() {
let guard = reservation_guard(MemoryGuardConfig {
limit_bytes: 1000,
pressure_threshold: f64::NAN,
cgroup_headroom: f64::NAN,
});
assert_eq!(guard.limit_bytes(), 1000);
guard.add_bytes(900);
assert!(guard.under_pressure());
}
#[test]
fn test_auto_detect_limit() {
let guard = MemoryGuard::new(MemoryGuardConfig::default());
assert!(
guard.limit_bytes() > 0,
"auto-detected limit should be positive"
);
}
}