use std::{
collections::BTreeMap,
fmt,
str::FromStr,
sync::{
Arc, Mutex, OnceLock, Weak,
atomic::{AtomicU64, Ordering},
},
};
use serde::{Deserialize, Serialize};
use sysinfo::{ProcessRefreshKind, ProcessesToUpdate, System, get_current_pid};
use thiserror::Error;
const AUTO_AVAILABLE_FRACTION: f64 = 0.80;
pub type MemoryResult<T> = Result<T, MemoryError>;
#[derive(Clone, Debug, Error, PartialEq)]
pub enum MemoryError {
#[error("invalid memory budget: {0}")]
InvalidBudget(String),
#[error("cannot resolve {budget} for {resource}: {basis} memory is unavailable")]
UnknownCapacity {
resource: String,
budget: MemoryBudget,
basis: &'static str,
},
#[error(
"memory budget exceeded for {resource}: requested {requested} bytes, \
{remaining} bytes remain"
)]
BudgetExceeded {
resource: String,
requested: u64,
remaining: u64,
},
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum CapacitySource {
OperatingSystem,
Cgroup,
Nvml,
Drm,
Dxgi,
Metal,
User,
Adaptive,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MemoryResourceKind {
Host,
Device,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum MemoryBudget {
#[default]
Auto,
Bytes(u64),
PercentTotal(f64),
PercentAvailable(f64),
}
impl MemoryBudget {
pub const fn bytes(bytes: u64) -> Self {
Self::Bytes(bytes)
}
pub fn percent_total(percent: f64) -> MemoryResult<Self> {
Ok(Self::PercentTotal(validate_percent(percent)? / 100.0))
}
pub fn percent_available(percent: f64) -> MemoryResult<Self> {
Ok(Self::PercentAvailable(validate_percent(percent)? / 100.0))
}
pub fn resolve(self, resource: &MemoryResource) -> MemoryResult<u64> {
let resolved = match self {
Self::Auto => resource
.available_bytes
.map(|bytes| scaled_bytes(bytes, AUTO_AVAILABLE_FRACTION))
.or(resource.total_bytes.map(|bytes| scaled_bytes(bytes, 0.5)))
.ok_or_else(|| MemoryError::UnknownCapacity {
resource: resource.name.clone(),
budget: self,
basis: "available",
})?,
Self::Bytes(bytes) => bytes,
Self::PercentTotal(fraction) => {
validate_fraction(fraction)?;
scaled_bytes(
resource
.total_bytes
.ok_or_else(|| MemoryError::UnknownCapacity {
resource: resource.name.clone(),
budget: self,
basis: "total",
})?,
fraction,
)
}
Self::PercentAvailable(fraction) => {
validate_fraction(fraction)?;
scaled_bytes(
resource
.available_bytes
.ok_or_else(|| MemoryError::UnknownCapacity {
resource: resource.name.clone(),
budget: self,
basis: "available",
})?,
fraction,
)
}
};
if resolved == 0 {
return Err(MemoryError::InvalidBudget(
"resolved budget must be greater than zero".into(),
));
}
Ok(resolved)
}
}
impl fmt::Display for MemoryBudget {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Auto => formatter.write_str("auto"),
Self::Bytes(bytes) => write!(formatter, "{bytes} B"),
Self::PercentTotal(value) => write!(formatter, "{}% total", value * 100.0),
Self::PercentAvailable(value) => {
write!(formatter, "{}% available", value * 100.0)
}
}
}
}
impl FromStr for MemoryBudget {
type Err = MemoryError;
fn from_str(input: &str) -> Result<Self, Self::Err> {
let normalized = input.trim().to_ascii_lowercase();
if normalized == "auto" {
return Ok(Self::Auto);
}
if let Some((percent, suffix)) = normalized.split_once('%') {
let percent = percent
.trim()
.parse::<f64>()
.map_err(|_| MemoryError::InvalidBudget(input.into()))?;
let suffix = suffix.trim();
return match suffix {
"" | "total" => Self::percent_total(percent),
"available" | "free" | "remaining" => Self::percent_available(percent),
_ => Err(MemoryError::InvalidBudget(input.into())),
};
}
parse_bytes(&normalized).map(Self::Bytes)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
pub struct MemoryPlan {
pub host: MemoryBudget,
pub device: Option<MemoryBudget>,
}
impl Default for MemoryPlan {
fn default() -> Self {
Self {
host: MemoryBudget::Auto,
device: Some(MemoryBudget::Auto),
}
}
}
impl MemoryPlan {
pub const fn host(host: MemoryBudget) -> Self {
Self { host, device: None }
}
pub const fn host_device(host: MemoryBudget, device: MemoryBudget) -> Self {
Self {
host,
device: Some(device),
}
}
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct DeviceIdentity {
pub adapter_index: usize,
pub vendor_id: u32,
pub device_id: u32,
pub pci_bus_id: String,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct MemoryResource {
pub id: String,
pub name: String,
pub kind: MemoryResourceKind,
pub total_bytes: Option<u64>,
pub available_bytes: Option<u64>,
pub capacity_source: CapacitySource,
pub device_identity: Option<DeviceIdentity>,
}
impl MemoryResource {
pub fn adaptive_device(id: impl Into<String>, name: impl Into<String>) -> Self {
Self {
id: id.into(),
name: name.into(),
kind: MemoryResourceKind::Device,
total_bytes: None,
available_bytes: None,
capacity_source: CapacitySource::Adaptive,
device_identity: None,
}
}
pub fn discover_device(
id: impl Into<String>,
name: impl Into<String>,
identity: DeviceIdentity,
fallback_bytes: u64,
) -> Self {
let mut resource = Self::adaptive_device(id, name);
resource.device_identity = Some(identity);
if let Some((total, available, source)) = refresh_device_memory(&resource) {
resource.total_bytes = Some(total);
resource.available_bytes = Some(available);
resource.capacity_source = source;
return resource;
}
resource.total_bytes = Some(fallback_bytes);
resource.available_bytes = Some(fallback_bytes);
resource
}
pub fn with_capacity(mut self, total_bytes: u64, available_bytes: Option<u64>) -> Self {
self.total_bytes = Some(total_bytes);
self.available_bytes = Some(available_bytes.unwrap_or(total_bytes).min(total_bytes));
self.capacity_source = CapacitySource::User;
self
}
pub const fn budget(&self, budget: MemoryBudget) -> MemoryBudget {
budget
}
}
#[derive(Debug)]
struct ResourceLedger {
snapshot: MemoryResource,
reserved: u64,
high_water: u64,
}
#[derive(Clone, Debug)]
pub struct MemoryState {
inner: Arc<MemoryStateInner>,
}
#[derive(Debug)]
struct MemoryStateInner {
resources: Mutex<BTreeMap<String, ResourceLedger>>,
process_high_water: AtomicU64,
}
impl MemoryState {
pub fn discover() -> Self {
let host = discover_host();
let mut resources = BTreeMap::new();
resources.insert(
host.id.clone(),
ResourceLedger {
snapshot: host,
reserved: 0,
high_water: 0,
},
);
Self {
inner: Arc::new(MemoryStateInner {
resources: Mutex::new(resources),
process_high_water: AtomicU64::new(0),
}),
}
}
pub fn current() -> Self {
static CURRENT: OnceLock<MemoryState> = OnceLock::new();
CURRENT.get_or_init(Self::discover).clone()
}
pub fn refresh(&self) {
self.refresh_inner();
}
fn refresh_inner(&self) -> Option<ProcessMemoryReport> {
let host = discover_host();
let process = self.sample_process_memory();
let mut resources = self
.inner
.resources
.lock()
.unwrap_or_else(|e| e.into_inner());
let ledger = resources
.entry(host.id.clone())
.or_insert_with(|| ResourceLedger {
snapshot: host.clone(),
reserved: 0,
high_water: 0,
});
ledger.snapshot = host;
for ledger in resources.values_mut() {
if ledger.snapshot.kind != MemoryResourceKind::Device
|| ledger.snapshot.capacity_source == CapacitySource::User
{
continue;
}
if let Some((total, available, source)) = refresh_device_memory(&ledger.snapshot) {
ledger.snapshot.total_bytes = Some(total);
ledger.snapshot.available_bytes = Some(available);
ledger.snapshot.capacity_source = source;
}
}
process
}
pub fn host(&self) -> MemoryResource {
self.resource("host").unwrap_or_else(discover_host)
}
pub fn register_device(&self, resource: MemoryResource) {
let mut resources = self
.inner
.resources
.lock()
.unwrap_or_else(|e| e.into_inner());
let ledger = resources
.entry(resource.id.clone())
.or_insert_with(|| ResourceLedger {
snapshot: resource.clone(),
reserved: 0,
high_water: 0,
});
if ledger.snapshot.capacity_source != CapacitySource::User
|| resource.capacity_source == CapacitySource::User
{
ledger.snapshot = resource;
}
}
pub fn resource(&self, id: &str) -> Option<MemoryResource> {
self.inner
.resources
.lock()
.unwrap_or_else(|e| e.into_inner())
.get(id)
.map(|ledger| ledger.snapshot.clone())
}
pub fn devices(&self) -> Vec<MemoryResource> {
self.inner
.resources
.lock()
.unwrap_or_else(|e| e.into_inner())
.values()
.filter(|ledger| ledger.snapshot.kind == MemoryResourceKind::Device)
.map(|ledger| ledger.snapshot.clone())
.collect()
}
pub fn pool(&self, resource_id: &str, budget: MemoryBudget) -> MemoryResult<MemoryPool> {
let resource = self.resource(resource_id).ok_or_else(|| {
MemoryError::InvalidBudget(format!("unknown memory resource {resource_id:?}"))
})?;
let capacity = budget.resolve(&resource)?;
Ok(MemoryPool {
inner: Arc::new(MemoryPoolInner {
state: Arc::downgrade(&self.inner),
resource_id: resource_id.to_owned(),
requested: budget,
capacity,
reserved: AtomicU64::new(0),
high_water: AtomicU64::new(0),
}),
})
}
pub fn report(&self) -> MemoryReport {
let process = self.refresh_inner();
let resources = self
.inner
.resources
.lock()
.unwrap_or_else(|e| e.into_inner());
MemoryReport {
process,
resources: resources
.values()
.map(|ledger| MemoryResourceReport {
resource: ledger.snapshot.clone(),
laddu_reserved_bytes: ledger.reserved,
laddu_high_water_bytes: ledger.high_water,
})
.collect(),
}
}
fn sample_process_memory(&self) -> Option<ProcessMemoryReport> {
let (resident_bytes, virtual_bytes) = discover_process_memory()?;
self.inner
.process_high_water
.fetch_max(resident_bytes, Ordering::AcqRel);
Some(ProcessMemoryReport {
resident_bytes,
virtual_bytes,
sampled_high_water_bytes: self.inner.process_high_water.load(Ordering::Acquire),
})
}
}
#[derive(Clone, Debug)]
pub struct MemoryPool {
inner: Arc<MemoryPoolInner>,
}
#[derive(Debug)]
struct MemoryPoolInner {
state: Weak<MemoryStateInner>,
resource_id: String,
requested: MemoryBudget,
capacity: u64,
reserved: AtomicU64,
high_water: AtomicU64,
}
impl MemoryPool {
pub fn requested(&self) -> MemoryBudget {
self.inner.requested
}
pub fn capacity(&self) -> u64 {
self.inner.capacity
}
pub fn reserved(&self) -> u64 {
self.inner.reserved.load(Ordering::Acquire)
}
pub fn remaining(&self) -> u64 {
self.capacity().saturating_sub(self.reserved())
}
pub fn high_water(&self) -> u64 {
self.inner.high_water.load(Ordering::Acquire)
}
pub fn reserve(&self, bytes: u64) -> MemoryResult<MemoryLease> {
let state = self.inner.state.upgrade();
let mut resources = state
.as_ref()
.map(|state| state.resources.lock().unwrap_or_else(|e| e.into_inner()));
let current = self.reserved();
let next = current
.checked_add(bytes)
.ok_or_else(|| budget_exceeded(self, bytes))?;
if next > self.capacity() {
return Err(budget_exceeded(self, bytes));
}
if let Some(resources) = resources.as_mut()
&& let Some(ledger) = resources.get_mut(&self.inner.resource_id)
{
let physical_limit = ledger
.snapshot
.available_bytes
.or(ledger.snapshot.total_bytes)
.unwrap_or(u64::MAX);
let shared_next = ledger.reserved.saturating_add(bytes);
if shared_next > physical_limit {
return Err(MemoryError::BudgetExceeded {
resource: ledger.snapshot.name.clone(),
requested: bytes,
remaining: physical_limit.saturating_sub(ledger.reserved),
});
}
ledger.reserved = shared_next;
ledger.high_water = ledger.high_water.max(shared_next);
}
self.inner.reserved.store(next, Ordering::Release);
update_max(&self.inner.high_water, next);
Ok(MemoryLease {
inner: Arc::new(MemoryLeaseInner {
pool: Arc::clone(&self.inner),
bytes,
}),
})
}
pub fn report(&self) -> MemoryPoolReport {
MemoryPoolReport {
resource_id: self.inner.resource_id.clone(),
requested: self.requested(),
effective_bytes: self.capacity(),
reserved_bytes: self.reserved(),
remaining_bytes: self.remaining(),
high_water_bytes: self.high_water(),
}
}
}
#[derive(Clone, Debug)]
pub struct MemoryLease {
inner: Arc<MemoryLeaseInner>,
}
#[derive(Debug)]
struct MemoryLeaseInner {
pool: Arc<MemoryPoolInner>,
bytes: u64,
}
impl MemoryLease {
pub fn bytes(&self) -> u64 {
self.inner.bytes
}
}
impl Drop for MemoryLeaseInner {
fn drop(&mut self) {
let pool = &self.pool;
pool.reserved.fetch_sub(self.bytes, Ordering::AcqRel);
if let Some(state) = pool.state.upgrade() {
let mut resources = state.resources.lock().unwrap_or_else(|e| e.into_inner());
if let Some(ledger) = resources.get_mut(&pool.resource_id) {
ledger.reserved = ledger.reserved.saturating_sub(self.bytes);
}
}
}
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct MemoryPoolReport {
pub resource_id: String,
pub requested: MemoryBudget,
pub effective_bytes: u64,
pub reserved_bytes: u64,
pub remaining_bytes: u64,
pub high_water_bytes: u64,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct MemoryResourceReport {
pub resource: MemoryResource,
pub laddu_reserved_bytes: u64,
pub laddu_high_water_bytes: u64,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct MemoryReport {
pub process: Option<ProcessMemoryReport>,
pub resources: Vec<MemoryResourceReport>,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct ProcessMemoryReport {
pub resident_bytes: u64,
pub virtual_bytes: u64,
pub sampled_high_water_bytes: u64,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct MemoryDecision {
pub label: String,
pub fixed_bytes: u64,
pub bytes_per_event: u64,
pub chunk_events: usize,
pub estimated_peak_bytes: u64,
pub actual_high_water_bytes: Option<u64>,
pub strategy: String,
}
impl MemoryDecision {
pub fn fit(
label: impl Into<String>,
fixed_bytes: u64,
bytes_per_event: u64,
available_bytes: u64,
event_limit: usize,
strategy: impl Into<String>,
) -> MemoryResult<Self> {
let label = label.into();
let per_event = bytes_per_event.max(1);
let capacity = available_bytes.saturating_sub(fixed_bytes);
let events = usize::try_from(capacity / per_event)
.unwrap_or(usize::MAX)
.min(event_limit);
if event_limit > 0 && events == 0 {
return Err(MemoryError::BudgetExceeded {
resource: label,
requested: fixed_bytes.saturating_add(per_event),
remaining: available_bytes,
});
}
let peak = fixed_bytes.saturating_add(per_event.saturating_mul(events as u64));
Ok(Self {
label,
fixed_bytes,
bytes_per_event,
chunk_events: events,
estimated_peak_bytes: peak,
actual_high_water_bytes: None,
strategy: strategy.into(),
})
}
}
fn discover_host() -> MemoryResource {
let mut system = System::new();
system.refresh_memory();
let mut total = system.total_memory();
let mut available = system.available_memory();
let mut capacity_source = CapacitySource::OperatingSystem;
if let Some((limit, available_in_group)) = discover_cgroup_memory()
&& limit < total
{
total = limit;
available = available.min(available_in_group);
capacity_source = CapacitySource::Cgroup;
}
MemoryResource {
id: "host".into(),
name: "Host memory".into(),
kind: MemoryResourceKind::Host,
total_bytes: Some(total),
available_bytes: Some(available),
capacity_source,
device_identity: None,
}
}
fn discover_cgroup_memory() -> Option<(u64, u64)> {
let pid = get_current_pid().ok()?;
let mut system = System::new();
system.refresh_processes_specifics(
ProcessesToUpdate::Some(&[pid]),
true,
ProcessRefreshKind::nothing().with_memory(),
);
let limits = system.process(pid)?.cgroup_limits()?;
Some((limits.total_memory, limits.free_memory))
}
fn discover_process_memory() -> Option<(u64, u64)> {
let pid = get_current_pid().ok()?;
let mut system = System::new();
system.refresh_processes_specifics(
ProcessesToUpdate::Some(&[pid]),
true,
ProcessRefreshKind::nothing().with_memory(),
);
let process = system.process(pid)?;
Some((process.memory(), process.virtual_memory()))
}
fn refresh_device_memory(resource: &MemoryResource) -> Option<(u64, u64, CapacitySource)> {
let identity = resource.device_identity.as_ref()?;
#[cfg(feature = "nvml")]
if let Some((total, available)) = discover_nvml_memory(&identity.pci_bus_id) {
return Some((total, available, CapacitySource::Nvml));
}
#[cfg(target_os = "windows")]
if let Some((total, available)) = discover_dxgi_memory(identity) {
return Some((total, available, CapacitySource::Dxgi));
}
#[cfg(target_os = "macos")]
if let Some((total, available)) = discover_metal_memory(identity, &resource.name) {
return Some((total, available, CapacitySource::Metal));
}
#[cfg(target_os = "linux")]
if let Some((total, available)) = discover_drm_memory(&identity.pci_bus_id) {
return Some((total, available, CapacitySource::Drm));
}
None
}
#[cfg(feature = "nvml")]
fn discover_nvml_memory(pci_bus_id: &str) -> Option<(u64, u64)> {
if pci_bus_id.is_empty() {
return None;
}
let nvml = nvml_wrapper::Nvml::init().ok()?;
let device = nvml.device_by_pci_bus_id(pci_bus_id).ok()?;
let memory = device.memory_info().ok()?;
Some((memory.total, memory.free))
}
#[cfg(target_os = "windows")]
fn discover_dxgi_memory(identity: &DeviceIdentity) -> Option<(u64, u64)> {
use windows::{
Win32::Graphics::Dxgi::{
CreateDXGIFactory1, DXGI_MEMORY_SEGMENT_GROUP_LOCAL, DXGI_QUERY_VIDEO_MEMORY_INFO,
IDXGIAdapter3, IDXGIFactory1,
},
core::Interface,
};
unsafe {
let factory: IDXGIFactory1 = CreateDXGIFactory1().ok()?;
let mut fallback = None;
for index in 0.. {
let Ok(adapter) = factory.EnumAdapters1(index) else {
break;
};
let Ok(description) = adapter.GetDesc1() else {
continue;
};
if description.VendorId != identity.vendor_id
|| description.DeviceId != identity.device_id
{
continue;
}
let adapter: IDXGIAdapter3 = adapter.cast().ok()?;
let mut memory = DXGI_QUERY_VIDEO_MEMORY_INFO::default();
adapter
.QueryVideoMemoryInfo(0, DXGI_MEMORY_SEGMENT_GROUP_LOCAL, &mut memory)
.ok()?;
let total = memory.Budget;
if total == 0 {
continue;
}
let snapshot = (total, total.saturating_sub(memory.CurrentUsage));
if index as usize == identity.adapter_index {
return Some(snapshot);
}
fallback.get_or_insert(snapshot);
}
return fallback;
}
}
#[cfg(target_os = "macos")]
fn discover_metal_memory(identity: &DeviceIdentity, expected_name: &str) -> Option<(u64, u64)> {
use objc2_metal::MTLDevice;
#[link(name = "CoreGraphics", kind = "framework")]
unsafe extern "C" {}
let devices = objc2_metal::MTLCopyAllDevices();
let device = (0..devices.count())
.map(|index| devices.objectAtIndex(index))
.find(|device| device.name().to_string() == expected_name)
.or_else(|| {
(identity.adapter_index < devices.count())
.then(|| devices.objectAtIndex(identity.adapter_index))
})?;
let total = device.recommendedMaxWorkingSetSize();
let used = device.currentAllocatedSize() as u64;
(total > 0).then_some((total, total.saturating_sub(used)))
}
#[cfg(target_os = "linux")]
fn discover_drm_memory(pci_bus_id: &str) -> Option<(u64, u64)> {
if pci_bus_id.is_empty() {
return None;
}
let entries = std::fs::read_dir("/sys/class/drm").ok()?;
for entry in entries.flatten() {
let name = entry.file_name();
if !name.to_string_lossy().starts_with("card") || name.to_string_lossy().contains('-') {
continue;
}
let device = entry.path().join("device");
let Ok(uevent) = std::fs::read_to_string(device.join("uevent")) else {
continue;
};
let matches_device = uevent.lines().any(|line| {
line.strip_prefix("PCI_SLOT_NAME=")
.is_some_and(|slot| slot.eq_ignore_ascii_case(pci_bus_id))
});
if !matches_device {
continue;
}
let Some(total) = read_sysfs_u64(device.join("mem_info_vram_total")) else {
continue;
};
let used = read_sysfs_u64(device.join("mem_info_vram_used")).unwrap_or(0);
return Some((total, total.saturating_sub(used)));
}
None
}
#[cfg(target_os = "linux")]
fn read_sysfs_u64(path: impl AsRef<std::path::Path>) -> Option<u64> {
std::fs::read_to_string(path).ok()?.trim().parse().ok()
}
fn validate_percent(percent: f64) -> MemoryResult<f64> {
if percent.is_finite() && percent > 0.0 && percent <= 100.0 {
Ok(percent)
} else {
Err(MemoryError::InvalidBudget(
"percentage must be finite and in (0, 100]".into(),
))
}
}
fn validate_fraction(fraction: f64) -> MemoryResult<()> {
validate_percent(fraction * 100.0).map(|_| ())
}
fn scaled_bytes(bytes: u64, fraction: f64) -> u64 {
((bytes as f64) * fraction).floor().min(u64::MAX as f64) as u64
}
fn parse_bytes(input: &str) -> MemoryResult<u64> {
let split = input
.find(|character: char| !character.is_ascii_digit() && character != '.')
.unwrap_or(input.len());
let (number, unit) = input.split_at(split);
let value = number
.trim()
.parse::<f64>()
.map_err(|_| MemoryError::InvalidBudget(input.into()))?;
if !value.is_finite() || value <= 0.0 {
return Err(MemoryError::InvalidBudget(input.into()));
}
let multiplier = match unit.trim() {
"" | "b" | "byte" | "bytes" => 1.0,
"kb" => 1_000.0,
"mb" => 1_000_000.0,
"gb" => 1_000_000_000.0,
"tb" => 1_000_000_000_000.0,
"kib" => 1024.0,
"mib" => 1024.0 * 1024.0,
"gib" => 1024.0 * 1024.0 * 1024.0,
"tib" => 1024.0 * 1024.0 * 1024.0 * 1024.0,
_ => return Err(MemoryError::InvalidBudget(input.into())),
};
let bytes = value * multiplier;
if bytes > u64::MAX as f64 {
return Err(MemoryError::InvalidBudget(input.into()));
}
Ok(bytes.floor() as u64)
}
fn budget_exceeded(pool: &MemoryPool, requested: u64) -> MemoryError {
MemoryError::BudgetExceeded {
resource: pool.inner.resource_id.clone(),
requested,
remaining: pool.remaining(),
}
}
fn update_max(value: &AtomicU64, candidate: u64) {
let mut current = value.load(Ordering::Acquire);
while candidate > current {
match value.compare_exchange_weak(current, candidate, Ordering::AcqRel, Ordering::Acquire) {
Ok(_) => break,
Err(observed) => current = observed,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn resource() -> MemoryResource {
MemoryResource {
id: "test".into(),
name: "Test".into(),
kind: MemoryResourceKind::Device,
total_bytes: Some(1_000),
available_bytes: Some(500),
capacity_source: CapacitySource::User,
device_identity: None,
}
}
#[test]
fn parses_absolute_and_percentage_budgets() {
assert_eq!(
"8 GiB".parse(),
Ok(MemoryBudget::Bytes(8 * 1024_u64.pow(3)))
);
assert_eq!("70% total".parse(), Ok(MemoryBudget::PercentTotal(0.7)));
assert_eq!(
"60% available".parse(),
Ok(MemoryBudget::PercentAvailable(0.6))
);
assert_eq!("auto".parse(), Ok(MemoryBudget::Auto));
}
#[test]
fn resolves_budgets_against_the_correct_capacity() {
let resource = resource();
assert_eq!(MemoryBudget::Auto.resolve(&resource), Ok(400));
assert_eq!(MemoryBudget::PercentTotal(0.5).resolve(&resource), Ok(500));
assert_eq!(
MemoryBudget::PercentAvailable(0.5).resolve(&resource),
Ok(250)
);
}
#[test]
fn leases_enforce_and_release_shared_capacity() {
let state = MemoryState::discover();
state.register_device(resource());
let pool = state.pool("test", MemoryBudget::Bytes(300)).unwrap();
let lease = pool.reserve(200).unwrap();
assert_eq!(pool.remaining(), 100);
assert!(pool.reserve(101).is_err());
drop(lease);
assert_eq!(pool.remaining(), 300);
assert_eq!(pool.high_water(), 200);
}
#[test]
fn decisions_fit_the_largest_safe_chunk() {
let decision = MemoryDecision::fit("test", 100, 8, 1_000, 1_000, "streaming").unwrap();
assert_eq!(decision.chunk_events, 112);
assert_eq!(decision.estimated_peak_bytes, 996);
}
#[test]
#[cfg(target_os = "linux")]
fn reports_current_process_memory() {
let state = MemoryState::discover();
let first = state.report().process.unwrap();
let second = state.report().process.unwrap();
assert!(first.resident_bytes > 0);
assert!(first.virtual_bytes >= first.resident_bytes);
assert!(second.sampled_high_water_bytes >= first.resident_bytes);
assert!(second.sampled_high_water_bytes >= second.resident_bytes);
}
}