Skip to main content

laddu_memory/
lib.rs

1//! Memory discovery, budgeting, reservation, and reporting for laddu.
2
3use std::{
4    collections::BTreeMap,
5    fmt,
6    str::FromStr,
7    sync::{
8        Arc, Mutex, OnceLock, Weak,
9        atomic::{AtomicU64, Ordering},
10    },
11};
12
13use serde::{Deserialize, Serialize};
14use sysinfo::{ProcessRefreshKind, ProcessesToUpdate, System, get_current_pid};
15use thiserror::Error;
16
17const AUTO_AVAILABLE_FRACTION: f64 = 0.80;
18
19/// Result type for memory planning operations.
20pub type MemoryResult<T> = Result<T, MemoryError>;
21
22/// Errors produced while discovering or reserving memory.
23#[derive(Clone, Debug, Error, PartialEq)]
24pub enum MemoryError {
25    /// A budget string or percentage is invalid.
26    #[error("invalid memory budget: {0}")]
27    InvalidBudget(String),
28    /// A percentage cannot be resolved because capacity telemetry is unavailable.
29    #[error("cannot resolve {budget} for {resource}: {basis} memory is unavailable")]
30    UnknownCapacity {
31        /// Resource label.
32        resource: String,
33        /// Requested budget.
34        budget: MemoryBudget,
35        /// Missing capacity basis.
36        basis: &'static str,
37    },
38    /// A reservation exceeds the effective pool limit.
39    #[error(
40        "memory budget exceeded for {resource}: requested {requested} bytes, \
41         {remaining} bytes remain"
42    )]
43    BudgetExceeded {
44        /// Resource label.
45        resource: String,
46        /// Requested reservation.
47        requested: u64,
48        /// Remaining reservable bytes.
49        remaining: u64,
50    },
51}
52
53/// The basis used to obtain a resource's capacity information.
54#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
55#[serde(rename_all = "snake_case")]
56pub enum CapacitySource {
57    /// Host operating-system telemetry.
58    OperatingSystem,
59    /// A process or container limit.
60    Cgroup,
61    /// NVIDIA Management Library telemetry.
62    Nvml,
63    /// Linux DRM/sysfs telemetry.
64    Drm,
65    /// Windows DXGI telemetry.
66    Dxgi,
67    /// Apple Metal working-set telemetry.
68    Metal,
69    /// Capacity supplied by the user.
70    User,
71    /// Capacity is not observable and planning is adaptive.
72    Adaptive,
73}
74
75/// Kind of physical memory resource.
76#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
77#[serde(rename_all = "snake_case")]
78pub enum MemoryResourceKind {
79    /// Host RAM.
80    Host,
81    /// Accelerator-local or unified memory.
82    Device,
83}
84
85/// A requested memory limit.
86#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
87#[serde(rename_all = "snake_case")]
88pub enum MemoryBudget {
89    /// Automatically use 80% of currently available memory.
90    #[default]
91    Auto,
92    /// An absolute number of bytes.
93    Bytes(u64),
94    /// A fraction in `(0, 1]` of total physical capacity.
95    PercentTotal(f64),
96    /// A fraction in `(0, 1]` of currently available capacity.
97    PercentAvailable(f64),
98}
99
100impl MemoryBudget {
101    /// Creates an absolute byte budget.
102    pub const fn bytes(bytes: u64) -> Self {
103        Self::Bytes(bytes)
104    }
105
106    /// Creates a percentage-of-total budget.
107    ///
108    /// # Errors
109    ///
110    /// Returns [`MemoryError::InvalidBudget`] unless `percent` is in `(0, 100]`.
111    pub fn percent_total(percent: f64) -> MemoryResult<Self> {
112        Ok(Self::PercentTotal(validate_percent(percent)? / 100.0))
113    }
114
115    /// Creates a percentage-of-available budget.
116    ///
117    /// # Errors
118    ///
119    /// Returns [`MemoryError::InvalidBudget`] unless `percent` is in `(0, 100]`.
120    pub fn percent_available(percent: f64) -> MemoryResult<Self> {
121        Ok(Self::PercentAvailable(validate_percent(percent)? / 100.0))
122    }
123
124    /// Resolves this request for a resource snapshot.
125    ///
126    /// # Errors
127    ///
128    /// Returns an error for zero budgets, invalid percentages, or unavailable
129    /// capacity telemetry.
130    pub fn resolve(self, resource: &MemoryResource) -> MemoryResult<u64> {
131        let resolved = match self {
132            Self::Auto => resource
133                .available_bytes
134                .map(|bytes| scaled_bytes(bytes, AUTO_AVAILABLE_FRACTION))
135                .or(resource.total_bytes.map(|bytes| scaled_bytes(bytes, 0.5)))
136                .ok_or_else(|| MemoryError::UnknownCapacity {
137                    resource: resource.name.clone(),
138                    budget: self,
139                    basis: "available",
140                })?,
141            Self::Bytes(bytes) => bytes,
142            Self::PercentTotal(fraction) => {
143                validate_fraction(fraction)?;
144                scaled_bytes(
145                    resource
146                        .total_bytes
147                        .ok_or_else(|| MemoryError::UnknownCapacity {
148                            resource: resource.name.clone(),
149                            budget: self,
150                            basis: "total",
151                        })?,
152                    fraction,
153                )
154            }
155            Self::PercentAvailable(fraction) => {
156                validate_fraction(fraction)?;
157                scaled_bytes(
158                    resource
159                        .available_bytes
160                        .ok_or_else(|| MemoryError::UnknownCapacity {
161                            resource: resource.name.clone(),
162                            budget: self,
163                            basis: "available",
164                        })?,
165                    fraction,
166                )
167            }
168        };
169        if resolved == 0 {
170            return Err(MemoryError::InvalidBudget(
171                "resolved budget must be greater than zero".into(),
172            ));
173        }
174        Ok(resolved)
175    }
176}
177
178impl fmt::Display for MemoryBudget {
179    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
180        match self {
181            Self::Auto => formatter.write_str("auto"),
182            Self::Bytes(bytes) => write!(formatter, "{bytes} B"),
183            Self::PercentTotal(value) => write!(formatter, "{}% total", value * 100.0),
184            Self::PercentAvailable(value) => {
185                write!(formatter, "{}% available", value * 100.0)
186            }
187        }
188    }
189}
190
191impl FromStr for MemoryBudget {
192    type Err = MemoryError;
193
194    fn from_str(input: &str) -> Result<Self, Self::Err> {
195        let normalized = input.trim().to_ascii_lowercase();
196        if normalized == "auto" {
197            return Ok(Self::Auto);
198        }
199        if let Some((percent, suffix)) = normalized.split_once('%') {
200            let percent = percent
201                .trim()
202                .parse::<f64>()
203                .map_err(|_| MemoryError::InvalidBudget(input.into()))?;
204            let suffix = suffix.trim();
205            return match suffix {
206                "" | "total" => Self::percent_total(percent),
207                "available" | "free" | "remaining" => Self::percent_available(percent),
208                _ => Err(MemoryError::InvalidBudget(input.into())),
209            };
210        }
211        parse_bytes(&normalized).map(Self::Bytes)
212    }
213}
214
215/// Host and optional accelerator budgets for one execution.
216#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
217pub struct MemoryPlan {
218    /// Host allocations, including source and staging buffers.
219    pub host: MemoryBudget,
220    /// Device allocations. Required only for accelerator execution.
221    pub device: Option<MemoryBudget>,
222}
223
224impl Default for MemoryPlan {
225    fn default() -> Self {
226        Self {
227            host: MemoryBudget::Auto,
228            device: Some(MemoryBudget::Auto),
229        }
230    }
231}
232
233impl MemoryPlan {
234    /// Creates a host-only plan.
235    pub const fn host(host: MemoryBudget) -> Self {
236        Self { host, device: None }
237    }
238
239    /// Creates a host-and-device plan.
240    pub const fn host_device(host: MemoryBudget, device: MemoryBudget) -> Self {
241        Self {
242            host,
243            device: Some(device),
244        }
245    }
246}
247
248/// Stable information used to match a runtime accelerator to platform telemetry.
249#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
250pub struct DeviceIdentity {
251    /// Backend adapter index.
252    pub adapter_index: usize,
253    /// PCI vendor identifier, or zero when unavailable.
254    pub vendor_id: u32,
255    /// PCI device identifier, or zero when unavailable.
256    pub device_id: u32,
257    /// PCI bus identifier, or an empty string when unavailable.
258    pub pci_bus_id: String,
259}
260
261/// Snapshot of one physical memory resource.
262#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
263pub struct MemoryResource {
264    /// Stable resource identifier.
265    pub id: String,
266    /// Human-readable resource name.
267    pub name: String,
268    /// Host or accelerator memory.
269    pub kind: MemoryResourceKind,
270    /// Total capacity when observable.
271    pub total_bytes: Option<u64>,
272    /// Currently available capacity when observable.
273    pub available_bytes: Option<u64>,
274    /// Source of capacity information.
275    pub capacity_source: CapacitySource,
276    /// Accelerator identity used to refresh platform telemetry.
277    pub device_identity: Option<DeviceIdentity>,
278}
279
280impl MemoryResource {
281    /// Creates an accelerator resource whose physical capacity is unavailable.
282    pub fn adaptive_device(id: impl Into<String>, name: impl Into<String>) -> Self {
283        Self {
284            id: id.into(),
285            name: name.into(),
286            kind: MemoryResourceKind::Device,
287            total_bytes: None,
288            available_bytes: None,
289            capacity_source: CapacitySource::Adaptive,
290            device_identity: None,
291        }
292    }
293
294    /// Discovers accelerator capacity from platform telemetry, falling back to
295    /// an adaptive backend limit when dedicated-memory telemetry is absent.
296    pub fn discover_device(
297        id: impl Into<String>,
298        name: impl Into<String>,
299        identity: DeviceIdentity,
300        fallback_bytes: u64,
301    ) -> Self {
302        let mut resource = Self::adaptive_device(id, name);
303        resource.device_identity = Some(identity);
304        if let Some((total, available, source)) = refresh_device_memory(&resource) {
305            resource.total_bytes = Some(total);
306            resource.available_bytes = Some(available);
307            resource.capacity_source = source;
308            return resource;
309        }
310        resource.total_bytes = Some(fallback_bytes);
311        resource.available_bytes = Some(fallback_bytes);
312        resource
313    }
314
315    /// Creates a resource with an explicit user-provided capacity.
316    pub fn with_capacity(mut self, total_bytes: u64, available_bytes: Option<u64>) -> Self {
317        self.total_bytes = Some(total_bytes);
318        self.available_bytes = Some(available_bytes.unwrap_or(total_bytes).min(total_bytes));
319        self.capacity_source = CapacitySource::User;
320        self
321    }
322
323    /// Creates a budget bound to this resource.
324    pub const fn budget(&self, budget: MemoryBudget) -> MemoryBudget {
325        budget
326    }
327}
328
329#[derive(Debug)]
330struct ResourceLedger {
331    snapshot: MemoryResource,
332    reserved: u64,
333    high_water: u64,
334}
335
336/// Live resource discovery and process-wide laddu reservation state.
337#[derive(Clone, Debug)]
338pub struct MemoryState {
339    inner: Arc<MemoryStateInner>,
340}
341
342#[derive(Debug)]
343struct MemoryStateInner {
344    resources: Mutex<BTreeMap<String, ResourceLedger>>,
345    process_high_water: AtomicU64,
346}
347
348impl MemoryState {
349    /// Discovers host memory and creates an independent reservation state.
350    pub fn discover() -> Self {
351        let host = discover_host();
352        let mut resources = BTreeMap::new();
353        resources.insert(
354            host.id.clone(),
355            ResourceLedger {
356                snapshot: host,
357                reserved: 0,
358                high_water: 0,
359            },
360        );
361        Self {
362            inner: Arc::new(MemoryStateInner {
363                resources: Mutex::new(resources),
364                process_high_water: AtomicU64::new(0),
365            }),
366        }
367    }
368
369    /// Returns the process-wide default state.
370    pub fn current() -> Self {
371        static CURRENT: OnceLock<MemoryState> = OnceLock::new();
372        CURRENT.get_or_init(Self::discover).clone()
373    }
374
375    /// Refreshes host total and available memory.
376    pub fn refresh(&self) {
377        self.refresh_inner();
378    }
379
380    fn refresh_inner(&self) -> Option<ProcessMemoryReport> {
381        let host = discover_host();
382        let process = self.sample_process_memory();
383        let mut resources = self
384            .inner
385            .resources
386            .lock()
387            .unwrap_or_else(|e| e.into_inner());
388        let ledger = resources
389            .entry(host.id.clone())
390            .or_insert_with(|| ResourceLedger {
391                snapshot: host.clone(),
392                reserved: 0,
393                high_water: 0,
394            });
395        ledger.snapshot = host;
396        for ledger in resources.values_mut() {
397            if ledger.snapshot.kind != MemoryResourceKind::Device
398                || ledger.snapshot.capacity_source == CapacitySource::User
399            {
400                continue;
401            }
402            if let Some((total, available, source)) = refresh_device_memory(&ledger.snapshot) {
403                ledger.snapshot.total_bytes = Some(total);
404                ledger.snapshot.available_bytes = Some(available);
405                ledger.snapshot.capacity_source = source;
406            }
407        }
408        process
409    }
410
411    /// Returns the current host snapshot.
412    pub fn host(&self) -> MemoryResource {
413        self.resource("host").unwrap_or_else(discover_host)
414    }
415
416    /// Registers or refreshes an accelerator resource.
417    pub fn register_device(&self, resource: MemoryResource) {
418        let mut resources = self
419            .inner
420            .resources
421            .lock()
422            .unwrap_or_else(|e| e.into_inner());
423        let ledger = resources
424            .entry(resource.id.clone())
425            .or_insert_with(|| ResourceLedger {
426                snapshot: resource.clone(),
427                reserved: 0,
428                high_water: 0,
429            });
430        // A user capacity override is authoritative until explicitly replaced
431        // by another user override.
432        if ledger.snapshot.capacity_source != CapacitySource::User
433            || resource.capacity_source == CapacitySource::User
434        {
435            ledger.snapshot = resource;
436        }
437    }
438
439    /// Returns one resource by stable identifier.
440    pub fn resource(&self, id: &str) -> Option<MemoryResource> {
441        self.inner
442            .resources
443            .lock()
444            .unwrap_or_else(|e| e.into_inner())
445            .get(id)
446            .map(|ledger| ledger.snapshot.clone())
447    }
448
449    /// Returns all registered accelerator resources.
450    pub fn devices(&self) -> Vec<MemoryResource> {
451        self.inner
452            .resources
453            .lock()
454            .unwrap_or_else(|e| e.into_inner())
455            .values()
456            .filter(|ledger| ledger.snapshot.kind == MemoryResourceKind::Device)
457            .map(|ledger| ledger.snapshot.clone())
458            .collect()
459    }
460
461    /// Resolves a budget and creates a local pool for a resource.
462    ///
463    /// # Errors
464    ///
465    /// Returns an error when the resource is unknown or the budget cannot be resolved.
466    pub fn pool(&self, resource_id: &str, budget: MemoryBudget) -> MemoryResult<MemoryPool> {
467        let resource = self.resource(resource_id).ok_or_else(|| {
468            MemoryError::InvalidBudget(format!("unknown memory resource {resource_id:?}"))
469        })?;
470        let capacity = budget.resolve(&resource)?;
471        Ok(MemoryPool {
472            inner: Arc::new(MemoryPoolInner {
473                state: Arc::downgrade(&self.inner),
474                resource_id: resource_id.to_owned(),
475                requested: budget,
476                capacity,
477                reserved: AtomicU64::new(0),
478                high_water: AtomicU64::new(0),
479            }),
480        })
481    }
482
483    /// Returns a structured report for all resources.
484    pub fn report(&self) -> MemoryReport {
485        let process = self.refresh_inner();
486        let resources = self
487            .inner
488            .resources
489            .lock()
490            .unwrap_or_else(|e| e.into_inner());
491        MemoryReport {
492            process,
493            resources: resources
494                .values()
495                .map(|ledger| MemoryResourceReport {
496                    resource: ledger.snapshot.clone(),
497                    laddu_reserved_bytes: ledger.reserved,
498                    laddu_high_water_bytes: ledger.high_water,
499                })
500                .collect(),
501        }
502    }
503
504    fn sample_process_memory(&self) -> Option<ProcessMemoryReport> {
505        let (resident_bytes, virtual_bytes) = discover_process_memory()?;
506        self.inner
507            .process_high_water
508            .fetch_max(resident_bytes, Ordering::AcqRel);
509        Some(ProcessMemoryReport {
510            resident_bytes,
511            virtual_bytes,
512            sampled_high_water_bytes: self.inner.process_high_water.load(Ordering::Acquire),
513        })
514    }
515}
516
517/// A resolved, reservable limit within one physical resource.
518#[derive(Clone, Debug)]
519pub struct MemoryPool {
520    inner: Arc<MemoryPoolInner>,
521}
522
523#[derive(Debug)]
524struct MemoryPoolInner {
525    state: Weak<MemoryStateInner>,
526    resource_id: String,
527    requested: MemoryBudget,
528    capacity: u64,
529    reserved: AtomicU64,
530    high_water: AtomicU64,
531}
532
533impl MemoryPool {
534    /// Requested budget specification.
535    pub fn requested(&self) -> MemoryBudget {
536        self.inner.requested
537    }
538
539    /// Resolved pool capacity in bytes.
540    pub fn capacity(&self) -> u64 {
541        self.inner.capacity
542    }
543
544    /// Currently reserved bytes.
545    pub fn reserved(&self) -> u64 {
546        self.inner.reserved.load(Ordering::Acquire)
547    }
548
549    /// Remaining reservable bytes.
550    pub fn remaining(&self) -> u64 {
551        self.capacity().saturating_sub(self.reserved())
552    }
553
554    /// Highest concurrent reservation observed by this pool.
555    pub fn high_water(&self) -> u64 {
556        self.inner.high_water.load(Ordering::Acquire)
557    }
558
559    /// Attempts to reserve `bytes` until the returned lease is dropped.
560    ///
561    /// # Errors
562    ///
563    /// Returns [`MemoryError::BudgetExceeded`] if the local or shared resource
564    /// limit lacks sufficient remaining capacity.
565    pub fn reserve(&self, bytes: u64) -> MemoryResult<MemoryLease> {
566        let state = self.inner.state.upgrade();
567        let mut resources = state
568            .as_ref()
569            .map(|state| state.resources.lock().unwrap_or_else(|e| e.into_inner()));
570
571        let current = self.reserved();
572        let next = current
573            .checked_add(bytes)
574            .ok_or_else(|| budget_exceeded(self, bytes))?;
575        if next > self.capacity() {
576            return Err(budget_exceeded(self, bytes));
577        }
578
579        if let Some(resources) = resources.as_mut()
580            && let Some(ledger) = resources.get_mut(&self.inner.resource_id)
581        {
582            let physical_limit = ledger
583                .snapshot
584                .available_bytes
585                .or(ledger.snapshot.total_bytes)
586                .unwrap_or(u64::MAX);
587            let shared_next = ledger.reserved.saturating_add(bytes);
588            if shared_next > physical_limit {
589                return Err(MemoryError::BudgetExceeded {
590                    resource: ledger.snapshot.name.clone(),
591                    requested: bytes,
592                    remaining: physical_limit.saturating_sub(ledger.reserved),
593                });
594            }
595            ledger.reserved = shared_next;
596            ledger.high_water = ledger.high_water.max(shared_next);
597        }
598
599        self.inner.reserved.store(next, Ordering::Release);
600        update_max(&self.inner.high_water, next);
601        Ok(MemoryLease {
602            inner: Arc::new(MemoryLeaseInner {
603                pool: Arc::clone(&self.inner),
604                bytes,
605            }),
606        })
607    }
608
609    /// Returns a snapshot of this pool's planning and usage.
610    pub fn report(&self) -> MemoryPoolReport {
611        MemoryPoolReport {
612            resource_id: self.inner.resource_id.clone(),
613            requested: self.requested(),
614            effective_bytes: self.capacity(),
615            reserved_bytes: self.reserved(),
616            remaining_bytes: self.remaining(),
617            high_water_bytes: self.high_water(),
618        }
619    }
620}
621
622/// A live memory reservation. Clones share one reservation.
623#[derive(Clone, Debug)]
624pub struct MemoryLease {
625    inner: Arc<MemoryLeaseInner>,
626}
627
628#[derive(Debug)]
629struct MemoryLeaseInner {
630    pool: Arc<MemoryPoolInner>,
631    bytes: u64,
632}
633
634impl MemoryLease {
635    /// Reserved byte count.
636    pub fn bytes(&self) -> u64 {
637        self.inner.bytes
638    }
639}
640
641impl Drop for MemoryLeaseInner {
642    fn drop(&mut self) {
643        let pool = &self.pool;
644        pool.reserved.fetch_sub(self.bytes, Ordering::AcqRel);
645        if let Some(state) = pool.state.upgrade() {
646            let mut resources = state.resources.lock().unwrap_or_else(|e| e.into_inner());
647            if let Some(ledger) = resources.get_mut(&pool.resource_id) {
648                ledger.reserved = ledger.reserved.saturating_sub(self.bytes);
649            }
650        }
651    }
652}
653
654/// Report for one resolved pool.
655#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
656pub struct MemoryPoolReport {
657    /// Stable resource identifier.
658    pub resource_id: String,
659    /// Requested budget.
660    pub requested: MemoryBudget,
661    /// Resolved capacity.
662    pub effective_bytes: u64,
663    /// Currently reserved bytes.
664    pub reserved_bytes: u64,
665    /// Remaining bytes.
666    pub remaining_bytes: u64,
667    /// Highest concurrent reservation.
668    pub high_water_bytes: u64,
669}
670
671/// Report for one physical resource.
672#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
673pub struct MemoryResourceReport {
674    /// Resource snapshot.
675    pub resource: MemoryResource,
676    /// Currently reserved by laddu.
677    pub laddu_reserved_bytes: u64,
678    /// Process-state high-water reservation.
679    pub laddu_high_water_bytes: u64,
680}
681
682/// Report covering all resources in a [`MemoryState`].
683#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
684pub struct MemoryReport {
685    /// Current process memory sampled when the report was generated.
686    pub process: Option<ProcessMemoryReport>,
687    /// Physical resource reports.
688    pub resources: Vec<MemoryResourceReport>,
689}
690
691/// Sampled operating-system memory counters for the current process.
692#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
693pub struct ProcessMemoryReport {
694    /// Current resident-set size.
695    pub resident_bytes: u64,
696    /// Current virtual-memory size.
697    pub virtual_bytes: u64,
698    /// Largest resident-set size observed by this [`MemoryState`].
699    ///
700    /// This is a sampled high-water mark, not an operating-system lifetime peak.
701    pub sampled_high_water_bytes: u64,
702}
703
704/// One memory-derived execution decision.
705#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
706pub struct MemoryDecision {
707    /// Operation or dataset label.
708    pub label: String,
709    /// Fixed bytes required regardless of event count.
710    pub fixed_bytes: u64,
711    /// Estimated incremental bytes per event.
712    pub bytes_per_event: u64,
713    /// Chosen internal event count.
714    pub chunk_events: usize,
715    /// Estimated peak tracked bytes.
716    pub estimated_peak_bytes: u64,
717    /// Actual tracked high-water bytes when known.
718    pub actual_high_water_bytes: Option<u64>,
719    /// Selected storage/execution strategy.
720    pub strategy: String,
721}
722
723impl MemoryDecision {
724    /// Derives the largest nonzero event chunk fitting within `available_bytes`.
725    ///
726    /// # Errors
727    ///
728    /// Returns [`MemoryError::BudgetExceeded`] when fixed cost plus one event
729    /// cannot fit.
730    pub fn fit(
731        label: impl Into<String>,
732        fixed_bytes: u64,
733        bytes_per_event: u64,
734        available_bytes: u64,
735        event_limit: usize,
736        strategy: impl Into<String>,
737    ) -> MemoryResult<Self> {
738        let label = label.into();
739        let per_event = bytes_per_event.max(1);
740        let capacity = available_bytes.saturating_sub(fixed_bytes);
741        let events = usize::try_from(capacity / per_event)
742            .unwrap_or(usize::MAX)
743            .min(event_limit);
744        if event_limit > 0 && events == 0 {
745            return Err(MemoryError::BudgetExceeded {
746                resource: label,
747                requested: fixed_bytes.saturating_add(per_event),
748                remaining: available_bytes,
749            });
750        }
751        let peak = fixed_bytes.saturating_add(per_event.saturating_mul(events as u64));
752        Ok(Self {
753            label,
754            fixed_bytes,
755            bytes_per_event,
756            chunk_events: events,
757            estimated_peak_bytes: peak,
758            actual_high_water_bytes: None,
759            strategy: strategy.into(),
760        })
761    }
762}
763
764fn discover_host() -> MemoryResource {
765    let mut system = System::new();
766    system.refresh_memory();
767    let mut total = system.total_memory();
768    let mut available = system.available_memory();
769    let mut capacity_source = CapacitySource::OperatingSystem;
770    if let Some((limit, available_in_group)) = discover_cgroup_memory()
771        && limit < total
772    {
773        total = limit;
774        available = available.min(available_in_group);
775        capacity_source = CapacitySource::Cgroup;
776    }
777    MemoryResource {
778        id: "host".into(),
779        name: "Host memory".into(),
780        kind: MemoryResourceKind::Host,
781        total_bytes: Some(total),
782        available_bytes: Some(available),
783        capacity_source,
784        device_identity: None,
785    }
786}
787
788fn discover_cgroup_memory() -> Option<(u64, u64)> {
789    let pid = get_current_pid().ok()?;
790    let mut system = System::new();
791    system.refresh_processes_specifics(
792        ProcessesToUpdate::Some(&[pid]),
793        true,
794        ProcessRefreshKind::nothing().with_memory(),
795    );
796    let limits = system.process(pid)?.cgroup_limits()?;
797    Some((limits.total_memory, limits.free_memory))
798}
799
800fn discover_process_memory() -> Option<(u64, u64)> {
801    let pid = get_current_pid().ok()?;
802    let mut system = System::new();
803    system.refresh_processes_specifics(
804        ProcessesToUpdate::Some(&[pid]),
805        true,
806        ProcessRefreshKind::nothing().with_memory(),
807    );
808    let process = system.process(pid)?;
809    Some((process.memory(), process.virtual_memory()))
810}
811
812fn refresh_device_memory(resource: &MemoryResource) -> Option<(u64, u64, CapacitySource)> {
813    let identity = resource.device_identity.as_ref()?;
814    #[cfg(feature = "nvml")]
815    if let Some((total, available)) = discover_nvml_memory(&identity.pci_bus_id) {
816        return Some((total, available, CapacitySource::Nvml));
817    }
818    #[cfg(target_os = "windows")]
819    if let Some((total, available)) = discover_dxgi_memory(identity) {
820        return Some((total, available, CapacitySource::Dxgi));
821    }
822    #[cfg(target_os = "macos")]
823    if let Some((total, available)) = discover_metal_memory(identity, &resource.name) {
824        return Some((total, available, CapacitySource::Metal));
825    }
826    #[cfg(target_os = "linux")]
827    if let Some((total, available)) = discover_drm_memory(&identity.pci_bus_id) {
828        return Some((total, available, CapacitySource::Drm));
829    }
830    None
831}
832
833#[cfg(feature = "nvml")]
834fn discover_nvml_memory(pci_bus_id: &str) -> Option<(u64, u64)> {
835    if pci_bus_id.is_empty() {
836        return None;
837    }
838    let nvml = nvml_wrapper::Nvml::init().ok()?;
839    let device = nvml.device_by_pci_bus_id(pci_bus_id).ok()?;
840    let memory = device.memory_info().ok()?;
841    Some((memory.total, memory.free))
842}
843
844#[cfg(target_os = "windows")]
845fn discover_dxgi_memory(identity: &DeviceIdentity) -> Option<(u64, u64)> {
846    use windows::{
847        Win32::Graphics::Dxgi::{
848            CreateDXGIFactory1, DXGI_MEMORY_SEGMENT_GROUP_LOCAL, DXGI_QUERY_VIDEO_MEMORY_INFO,
849            IDXGIAdapter3, IDXGIFactory1,
850        },
851        core::Interface,
852    };
853
854    // SAFETY: DXGI factory and adapter methods own their returned COM interfaces,
855    // and all output pointers are provided by the windows crate.
856    unsafe {
857        let factory: IDXGIFactory1 = CreateDXGIFactory1().ok()?;
858        let mut fallback = None;
859        for index in 0.. {
860            let Ok(adapter) = factory.EnumAdapters1(index) else {
861                break;
862            };
863            let Ok(description) = adapter.GetDesc1() else {
864                continue;
865            };
866            if description.VendorId != identity.vendor_id
867                || description.DeviceId != identity.device_id
868            {
869                continue;
870            }
871            let adapter: IDXGIAdapter3 = adapter.cast().ok()?;
872            let mut memory = DXGI_QUERY_VIDEO_MEMORY_INFO::default();
873            adapter
874                .QueryVideoMemoryInfo(0, DXGI_MEMORY_SEGMENT_GROUP_LOCAL, &mut memory)
875                .ok()?;
876            let total = memory.Budget;
877            if total == 0 {
878                continue;
879            }
880            let snapshot = (total, total.saturating_sub(memory.CurrentUsage));
881            if index as usize == identity.adapter_index {
882                return Some(snapshot);
883            }
884            fallback.get_or_insert(snapshot);
885        }
886        return fallback;
887    }
888}
889
890#[cfg(target_os = "macos")]
891fn discover_metal_memory(identity: &DeviceIdentity, expected_name: &str) -> Option<(u64, u64)> {
892    use objc2_metal::MTLDevice;
893
894    #[link(name = "CoreGraphics", kind = "framework")]
895    unsafe extern "C" {}
896
897    let devices = objc2_metal::MTLCopyAllDevices();
898    let device = (0..devices.count())
899        .map(|index| devices.objectAtIndex(index))
900        .find(|device| device.name().to_string() == expected_name)
901        .or_else(|| {
902            (identity.adapter_index < devices.count())
903                .then(|| devices.objectAtIndex(identity.adapter_index))
904        })?;
905    let total = device.recommendedMaxWorkingSetSize();
906    let used = device.currentAllocatedSize() as u64;
907    (total > 0).then_some((total, total.saturating_sub(used)))
908}
909
910#[cfg(target_os = "linux")]
911fn discover_drm_memory(pci_bus_id: &str) -> Option<(u64, u64)> {
912    if pci_bus_id.is_empty() {
913        return None;
914    }
915    let entries = std::fs::read_dir("/sys/class/drm").ok()?;
916    for entry in entries.flatten() {
917        let name = entry.file_name();
918        if !name.to_string_lossy().starts_with("card") || name.to_string_lossy().contains('-') {
919            continue;
920        }
921        let device = entry.path().join("device");
922        let Ok(uevent) = std::fs::read_to_string(device.join("uevent")) else {
923            continue;
924        };
925        let matches_device = uevent.lines().any(|line| {
926            line.strip_prefix("PCI_SLOT_NAME=")
927                .is_some_and(|slot| slot.eq_ignore_ascii_case(pci_bus_id))
928        });
929        if !matches_device {
930            continue;
931        }
932        let Some(total) = read_sysfs_u64(device.join("mem_info_vram_total")) else {
933            continue;
934        };
935        let used = read_sysfs_u64(device.join("mem_info_vram_used")).unwrap_or(0);
936        return Some((total, total.saturating_sub(used)));
937    }
938    None
939}
940
941#[cfg(target_os = "linux")]
942fn read_sysfs_u64(path: impl AsRef<std::path::Path>) -> Option<u64> {
943    std::fs::read_to_string(path).ok()?.trim().parse().ok()
944}
945
946fn validate_percent(percent: f64) -> MemoryResult<f64> {
947    if percent.is_finite() && percent > 0.0 && percent <= 100.0 {
948        Ok(percent)
949    } else {
950        Err(MemoryError::InvalidBudget(
951            "percentage must be finite and in (0, 100]".into(),
952        ))
953    }
954}
955
956fn validate_fraction(fraction: f64) -> MemoryResult<()> {
957    validate_percent(fraction * 100.0).map(|_| ())
958}
959
960fn scaled_bytes(bytes: u64, fraction: f64) -> u64 {
961    ((bytes as f64) * fraction).floor().min(u64::MAX as f64) as u64
962}
963
964fn parse_bytes(input: &str) -> MemoryResult<u64> {
965    let split = input
966        .find(|character: char| !character.is_ascii_digit() && character != '.')
967        .unwrap_or(input.len());
968    let (number, unit) = input.split_at(split);
969    let value = number
970        .trim()
971        .parse::<f64>()
972        .map_err(|_| MemoryError::InvalidBudget(input.into()))?;
973    if !value.is_finite() || value <= 0.0 {
974        return Err(MemoryError::InvalidBudget(input.into()));
975    }
976    let multiplier = match unit.trim() {
977        "" | "b" | "byte" | "bytes" => 1.0,
978        "kb" => 1_000.0,
979        "mb" => 1_000_000.0,
980        "gb" => 1_000_000_000.0,
981        "tb" => 1_000_000_000_000.0,
982        "kib" => 1024.0,
983        "mib" => 1024.0 * 1024.0,
984        "gib" => 1024.0 * 1024.0 * 1024.0,
985        "tib" => 1024.0 * 1024.0 * 1024.0 * 1024.0,
986        _ => return Err(MemoryError::InvalidBudget(input.into())),
987    };
988    let bytes = value * multiplier;
989    if bytes > u64::MAX as f64 {
990        return Err(MemoryError::InvalidBudget(input.into()));
991    }
992    Ok(bytes.floor() as u64)
993}
994
995fn budget_exceeded(pool: &MemoryPool, requested: u64) -> MemoryError {
996    MemoryError::BudgetExceeded {
997        resource: pool.inner.resource_id.clone(),
998        requested,
999        remaining: pool.remaining(),
1000    }
1001}
1002
1003fn update_max(value: &AtomicU64, candidate: u64) {
1004    let mut current = value.load(Ordering::Acquire);
1005    while candidate > current {
1006        match value.compare_exchange_weak(current, candidate, Ordering::AcqRel, Ordering::Acquire) {
1007            Ok(_) => break,
1008            Err(observed) => current = observed,
1009        }
1010    }
1011}
1012
1013#[cfg(test)]
1014mod tests {
1015    use super::*;
1016
1017    fn resource() -> MemoryResource {
1018        MemoryResource {
1019            id: "test".into(),
1020            name: "Test".into(),
1021            kind: MemoryResourceKind::Device,
1022            total_bytes: Some(1_000),
1023            available_bytes: Some(500),
1024            capacity_source: CapacitySource::User,
1025            device_identity: None,
1026        }
1027    }
1028
1029    #[test]
1030    fn parses_absolute_and_percentage_budgets() {
1031        assert_eq!(
1032            "8 GiB".parse(),
1033            Ok(MemoryBudget::Bytes(8 * 1024_u64.pow(3)))
1034        );
1035        assert_eq!("70% total".parse(), Ok(MemoryBudget::PercentTotal(0.7)));
1036        assert_eq!(
1037            "60% available".parse(),
1038            Ok(MemoryBudget::PercentAvailable(0.6))
1039        );
1040        assert_eq!("auto".parse(), Ok(MemoryBudget::Auto));
1041    }
1042
1043    #[test]
1044    fn resolves_budgets_against_the_correct_capacity() {
1045        let resource = resource();
1046        assert_eq!(MemoryBudget::Auto.resolve(&resource), Ok(400));
1047        assert_eq!(MemoryBudget::PercentTotal(0.5).resolve(&resource), Ok(500));
1048        assert_eq!(
1049            MemoryBudget::PercentAvailable(0.5).resolve(&resource),
1050            Ok(250)
1051        );
1052    }
1053
1054    #[test]
1055    fn leases_enforce_and_release_shared_capacity() {
1056        let state = MemoryState::discover();
1057        state.register_device(resource());
1058        let pool = state.pool("test", MemoryBudget::Bytes(300)).unwrap();
1059        let lease = pool.reserve(200).unwrap();
1060        assert_eq!(pool.remaining(), 100);
1061        assert!(pool.reserve(101).is_err());
1062        drop(lease);
1063        assert_eq!(pool.remaining(), 300);
1064        assert_eq!(pool.high_water(), 200);
1065    }
1066
1067    #[test]
1068    fn decisions_fit_the_largest_safe_chunk() {
1069        let decision = MemoryDecision::fit("test", 100, 8, 1_000, 1_000, "streaming").unwrap();
1070        assert_eq!(decision.chunk_events, 112);
1071        assert_eq!(decision.estimated_peak_bytes, 996);
1072    }
1073
1074    #[test]
1075    #[cfg(target_os = "linux")]
1076    fn reports_current_process_memory() {
1077        let state = MemoryState::discover();
1078        let first = state.report().process.unwrap();
1079        let second = state.report().process.unwrap();
1080        assert!(first.resident_bytes > 0);
1081        assert!(first.virtual_bytes >= first.resident_bytes);
1082        assert!(second.sampled_high_water_bytes >= first.resident_bytes);
1083        assert!(second.sampled_high_water_bytes >= second.resident_bytes);
1084    }
1085}