#[cfg(any(target_os = "linux", windows))]
use std::collections::HashMap;
#[cfg(any(target_os = "macos", target_os = "linux", windows))]
use std::collections::HashSet;
const MEMORY_LIMIT_FRACTION: f64 = 0.8;
#[must_use]
pub(super) fn default_limit() -> Option<u64> {
#[expect(
clippy::cast_precision_loss,
clippy::cast_possible_truncation,
clippy::cast_sign_loss
)]
let limit = (physical_memory()? as f64 * MEMORY_LIMIT_FRACTION) as u64;
Some(limit)
}
#[must_use]
pub(super) fn gib(bytes: u64) -> String {
const GIB: u64 = 1 << 30;
let whole = bytes / GIB;
let tenth = bytes % GIB * 10 / GIB;
format!("{whole}.{tenth} GiB")
}
#[must_use]
pub(super) fn exceeded_failure(used: u64, limit: u64) -> String {
let total = physical_memory().unwrap_or(limit);
format!(
"terminated: exceeded memory limit (used ~{} of {})",
gib(used),
gib(total),
)
}
#[cfg(target_os = "macos")]
#[must_use]
pub(super) fn physical_memory() -> Option<u64> {
let mut size = 0_u64;
let mut len = std::mem::size_of::<u64>();
let ok = unsafe {
libc::sysctlbyname(
c"hw.memsize".as_ptr(),
std::ptr::addr_of_mut!(size).cast(),
std::ptr::addr_of_mut!(len),
std::ptr::null_mut(),
0,
)
};
(ok == 0).then_some(size)
}
#[cfg(target_os = "linux")]
#[must_use]
pub(super) fn physical_memory() -> Option<u64> {
let pages = unsafe { libc::sysconf(libc::_SC_PHYS_PAGES) };
let page_size = unsafe { libc::sysconf(libc::_SC_PAGESIZE) };
if pages <= 0 || page_size <= 0 {
return None;
}
Some(
u64::try_from(pages)
.ok()?
.saturating_mul(u64::try_from(page_size).ok()?),
)
}
#[cfg(windows)]
#[must_use]
pub(super) fn physical_memory() -> Option<u64> {
use windows_sys::Win32::System::SystemInformation::{GlobalMemoryStatusEx, MEMORYSTATUSEX};
let mut status: MEMORYSTATUSEX = unsafe { std::mem::zeroed() };
status.dwLength =
u32::try_from(std::mem::size_of::<MEMORYSTATUSEX>()).expect("memory status fits in u32");
let ok = unsafe { GlobalMemoryStatusEx(std::ptr::addr_of_mut!(status)) };
(ok != 0).then_some(status.ullTotalPhys)
}
#[cfg(not(any(target_os = "macos", target_os = "linux", windows)))]
#[must_use]
pub(super) fn physical_memory() -> Option<u64> {
None
}
#[cfg(target_os = "macos")]
#[must_use]
pub(super) fn tree_rss(root: u32) -> Option<u64> {
let mut total = process_rss(root)?;
let mut seen = HashSet::from([root]);
let mut queue = vec![root];
while let Some(pid) = queue.pop() {
for child in children_of(pid) {
if !seen.insert(child) {
continue;
}
if let Some(rss) = process_rss(child) {
total = total.saturating_add(rss);
}
queue.push(child);
}
}
Some(total)
}
#[cfg(target_os = "linux")]
#[must_use]
pub(super) fn tree_rss(root: u32) -> Option<u64> {
sum_descendants(root, &children_map())
}
#[cfg(windows)]
#[must_use]
pub(super) fn tree_rss(root: u32) -> Option<u64> {
sum_descendants(root, &children_map()?)
}
#[cfg(not(any(target_os = "macos", target_os = "linux", windows)))]
#[must_use]
pub(super) fn tree_rss(_root: u32) -> Option<u64> {
None
}
#[cfg(any(target_os = "linux", windows))]
fn sum_descendants(root: u32, children: &HashMap<u32, Vec<u32>>) -> Option<u64> {
let mut total = process_rss(root)?;
let mut seen = HashSet::from([root]);
let mut queue = vec![root];
while let Some(pid) = queue.pop() {
for &child in children.get(&pid).into_iter().flatten() {
if !seen.insert(child) {
continue;
}
if let Some(rss) = process_rss(child) {
total = total.saturating_add(rss);
}
queue.push(child);
}
}
Some(total)
}
#[cfg(target_os = "macos")]
fn process_rss(pid: u32) -> Option<u64> {
let id = libc::c_int::try_from(pid).ok()?;
let mut info: libc::rusage_info_v2 = unsafe { std::mem::zeroed() };
let ok = unsafe {
libc::proc_pid_rusage(
id,
libc::RUSAGE_INFO_V2,
std::ptr::addr_of_mut!(info).cast(),
)
};
(ok == 0).then_some(info.ri_phys_footprint)
}
#[cfg(target_os = "macos")]
fn children_of(pid: u32) -> Vec<u32> {
const CHILD_CAPACITY: usize = 1024;
let Ok(parent) = libc::pid_t::try_from(pid) else {
return Vec::new();
};
let mut buffer: Vec<libc::pid_t> = vec![0; CHILD_CAPACITY];
let bytes = libc::c_int::try_from(std::mem::size_of_val(buffer.as_slice()))
.expect("the child buffer fits in a C int");
let written = unsafe { libc::proc_listchildpids(parent, buffer.as_mut_ptr().cast(), bytes) };
let Ok(filled) = usize::try_from(written) else {
return Vec::new();
};
buffer.truncate(filled.min(CHILD_CAPACITY));
buffer
.into_iter()
.filter_map(|p| u32::try_from(p).ok())
.collect()
}
#[cfg(target_os = "linux")]
fn process_rss(pid: u32) -> Option<u64> {
let statm = std::fs::read_to_string(format!("/proc/{pid}/statm")).ok()?;
let mut fields = statm.split_whitespace();
fields.next()?; let resident_pages: u64 = fields.next()?.parse().ok()?;
Some(resident_pages.saturating_mul(page_size()))
}
#[cfg(target_os = "linux")]
fn page_size() -> u64 {
let size = unsafe { libc::sysconf(libc::_SC_PAGESIZE) };
u64::try_from(size).unwrap_or(4096)
}
#[cfg(windows)]
fn process_rss(pid: u32) -> Option<u64> {
use windows_sys::Win32::Foundation::CloseHandle;
use windows_sys::Win32::System::ProcessStatus::{
GetProcessMemoryInfo, PROCESS_MEMORY_COUNTERS,
};
use windows_sys::Win32::System::Threading::{
OpenProcess, PROCESS_QUERY_LIMITED_INFORMATION, PROCESS_VM_READ,
};
let process =
unsafe { OpenProcess(PROCESS_QUERY_LIMITED_INFORMATION | PROCESS_VM_READ, 0, pid) };
if process == 0 {
return None;
}
let mut counters: PROCESS_MEMORY_COUNTERS = unsafe { std::mem::zeroed() };
counters.cb = u32::try_from(std::mem::size_of::<PROCESS_MEMORY_COUNTERS>())
.expect("memory counters fit in u32");
let ok =
unsafe { GetProcessMemoryInfo(process, std::ptr::addr_of_mut!(counters), counters.cb) };
unsafe { CloseHandle(process) };
if ok == 0 {
return None;
}
u64::try_from(counters.WorkingSetSize).ok()
}
#[cfg(target_os = "linux")]
fn children_map() -> HashMap<u32, Vec<u32>> {
let Ok(entries) = std::fs::read_dir("/proc") else {
return HashMap::new();
};
let mut children: HashMap<u32, Vec<u32>> = HashMap::new();
for entry in entries.flatten() {
let Some(pid) = entry
.file_name()
.to_str()
.and_then(|n| n.parse::<u32>().ok())
else {
continue;
};
if let Some(parent) = parent_of(pid) {
children.entry(parent).or_default().push(pid);
}
}
children
}
#[cfg(target_os = "linux")]
fn parent_of(pid: u32) -> Option<u32> {
let stat = std::fs::read_to_string(format!("/proc/{pid}/stat")).ok()?;
let rest = stat.get(stat.rfind(')')? + 1..)?;
let mut fields = rest.split_whitespace();
fields.next()?; fields.next()?.parse().ok()
}
#[cfg(windows)]
fn children_map() -> Option<HashMap<u32, Vec<u32>>> {
use windows_sys::Win32::Foundation::{CloseHandle, INVALID_HANDLE_VALUE};
use windows_sys::Win32::System::Diagnostics::ToolHelp::{
CreateToolhelp32Snapshot, PROCESSENTRY32W, Process32FirstW, Process32NextW,
TH32CS_SNAPPROCESS,
};
let snapshot = unsafe { CreateToolhelp32Snapshot(TH32CS_SNAPPROCESS, 0) };
if snapshot == INVALID_HANDLE_VALUE {
return None;
}
let mut entry: PROCESSENTRY32W = unsafe { std::mem::zeroed() };
entry.dwSize =
u32::try_from(std::mem::size_of::<PROCESSENTRY32W>()).expect("process entry fits in u32");
let mut children: HashMap<u32, Vec<u32>> = HashMap::new();
let mut more = unsafe { Process32FirstW(snapshot, std::ptr::addr_of_mut!(entry)) };
while more != 0 {
children
.entry(entry.th32ParentProcessID)
.or_default()
.push(entry.th32ProcessID);
more = unsafe { Process32NextW(snapshot, std::ptr::addr_of_mut!(entry)) };
}
unsafe { CloseHandle(snapshot) };
Some(children)
}
#[cfg(test)]
mod tests {
use super::{default_limit, exceeded_failure, gib, physical_memory, tree_rss};
#[test]
fn gib_renders_one_decimal_place() {
assert_eq!(gib(0), "0.0 GiB");
assert_eq!(gib(1 << 30), "1.0 GiB");
assert_eq!(gib(3 * (1 << 30) / 2), "1.5 GiB");
assert_eq!(gib(15 * (1 << 30) / 2 + 1), "7.5 GiB");
}
#[test]
fn the_memory_failure_names_the_amount_against_the_machines_total() {
let line = exceeded_failure(0, 1 << 30);
assert!(
line.starts_with("terminated: exceeded memory limit (used ~0.0 GiB of "),
"line: {line}"
);
}
#[test]
fn the_default_limit_is_eighty_percent_of_physical_memory() {
let Some(total) = physical_memory() else {
return;
};
let limit = default_limit().expect("a known total has a limit");
assert_eq!(limit, total.saturating_mul(4) / 5);
}
#[test]
fn the_current_process_tree_has_a_nonzero_footprint() {
let rss = tree_rss(std::process::id()).expect("the measuring process must be readable");
assert!(rss > 0, "footprint should be non-zero, got {rss}");
}
}