1use 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
19pub type MemoryResult<T> = Result<T, MemoryError>;
21
22#[derive(Clone, Debug, Error, PartialEq)]
24pub enum MemoryError {
25 #[error("invalid memory budget: {0}")]
27 InvalidBudget(String),
28 #[error("cannot resolve {budget} for {resource}: {basis} memory is unavailable")]
30 UnknownCapacity {
31 resource: String,
33 budget: MemoryBudget,
35 basis: &'static str,
37 },
38 #[error(
40 "memory budget exceeded for {resource}: requested {requested} bytes, \
41 {remaining} bytes remain"
42 )]
43 BudgetExceeded {
44 resource: String,
46 requested: u64,
48 remaining: u64,
50 },
51}
52
53#[derive(Clone, Copy, Debug, PartialEq, Eq, Serialize, Deserialize)]
55#[serde(rename_all = "snake_case")]
56pub enum CapacitySource {
57 OperatingSystem,
59 Cgroup,
61 Nvml,
63 Drm,
65 Dxgi,
67 Metal,
69 User,
71 Adaptive,
73}
74
75#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
77#[serde(rename_all = "snake_case")]
78pub enum MemoryResourceKind {
79 Host,
81 Device,
83}
84
85#[derive(Clone, Copy, Debug, Default, PartialEq, Serialize, Deserialize)]
87#[serde(rename_all = "snake_case")]
88pub enum MemoryBudget {
89 #[default]
91 Auto,
92 Bytes(u64),
94 PercentTotal(f64),
96 PercentAvailable(f64),
98}
99
100impl MemoryBudget {
101 pub const fn bytes(bytes: u64) -> Self {
103 Self::Bytes(bytes)
104 }
105
106 pub fn percent_total(percent: f64) -> MemoryResult<Self> {
112 Ok(Self::PercentTotal(validate_percent(percent)? / 100.0))
113 }
114
115 pub fn percent_available(percent: f64) -> MemoryResult<Self> {
121 Ok(Self::PercentAvailable(validate_percent(percent)? / 100.0))
122 }
123
124 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#[derive(Clone, Copy, Debug, PartialEq, Serialize, Deserialize)]
217pub struct MemoryPlan {
218 pub host: MemoryBudget,
220 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 pub const fn host(host: MemoryBudget) -> Self {
236 Self { host, device: None }
237 }
238
239 pub const fn host_device(host: MemoryBudget, device: MemoryBudget) -> Self {
241 Self {
242 host,
243 device: Some(device),
244 }
245 }
246}
247
248#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
250pub struct DeviceIdentity {
251 pub adapter_index: usize,
253 pub vendor_id: u32,
255 pub device_id: u32,
257 pub pci_bus_id: String,
259}
260
261#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
263pub struct MemoryResource {
264 pub id: String,
266 pub name: String,
268 pub kind: MemoryResourceKind,
270 pub total_bytes: Option<u64>,
272 pub available_bytes: Option<u64>,
274 pub capacity_source: CapacitySource,
276 pub device_identity: Option<DeviceIdentity>,
278}
279
280impl MemoryResource {
281 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 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 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 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#[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 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 pub fn current() -> Self {
371 static CURRENT: OnceLock<MemoryState> = OnceLock::new();
372 CURRENT.get_or_init(Self::discover).clone()
373 }
374
375 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 pub fn host(&self) -> MemoryResource {
413 self.resource("host").unwrap_or_else(discover_host)
414 }
415
416 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 if ledger.snapshot.capacity_source != CapacitySource::User
433 || resource.capacity_source == CapacitySource::User
434 {
435 ledger.snapshot = resource;
436 }
437 }
438
439 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 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 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 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#[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 pub fn requested(&self) -> MemoryBudget {
536 self.inner.requested
537 }
538
539 pub fn capacity(&self) -> u64 {
541 self.inner.capacity
542 }
543
544 pub fn reserved(&self) -> u64 {
546 self.inner.reserved.load(Ordering::Acquire)
547 }
548
549 pub fn remaining(&self) -> u64 {
551 self.capacity().saturating_sub(self.reserved())
552 }
553
554 pub fn high_water(&self) -> u64 {
556 self.inner.high_water.load(Ordering::Acquire)
557 }
558
559 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 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#[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 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#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
656pub struct MemoryPoolReport {
657 pub resource_id: String,
659 pub requested: MemoryBudget,
661 pub effective_bytes: u64,
663 pub reserved_bytes: u64,
665 pub remaining_bytes: u64,
667 pub high_water_bytes: u64,
669}
670
671#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
673pub struct MemoryResourceReport {
674 pub resource: MemoryResource,
676 pub laddu_reserved_bytes: u64,
678 pub laddu_high_water_bytes: u64,
680}
681
682#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
684pub struct MemoryReport {
685 pub process: Option<ProcessMemoryReport>,
687 pub resources: Vec<MemoryResourceReport>,
689}
690
691#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
693pub struct ProcessMemoryReport {
694 pub resident_bytes: u64,
696 pub virtual_bytes: u64,
698 pub sampled_high_water_bytes: u64,
702}
703
704#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
706pub struct MemoryDecision {
707 pub label: String,
709 pub fixed_bytes: u64,
711 pub bytes_per_event: u64,
713 pub chunk_events: usize,
715 pub estimated_peak_bytes: u64,
717 pub actual_high_water_bytes: Option<u64>,
719 pub strategy: String,
721}
722
723impl MemoryDecision {
724 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 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}