Skip to main content

laddu_memory/
pool.rs

1use std::sync::{
2    Arc, Weak,
3    atomic::{AtomicU64, Ordering},
4};
5
6use crate::{
7    budget::MemoryBudget,
8    discovery::CapacitySnapshot,
9    error::{MemoryError, MemoryResult},
10    report::MemoryPoolReport,
11    resource::{CapacitySource, MemoryResource},
12    state::MemoryStateInner,
13};
14
15#[derive(Debug)]
16pub(crate) struct ResourceLedger {
17    pub(crate) snapshot: MemoryResource,
18    pub(crate) reserved: u64,
19    pub(crate) high_water: u64,
20}
21
22impl ResourceLedger {
23    pub(crate) fn new(mut snapshot: MemoryResource) -> Self {
24        snapshot.normalize_capacity();
25        debug_assert_eq!(snapshot.validate(), Ok(()));
26        Self {
27            snapshot,
28            reserved: 0,
29            high_water: 0,
30        }
31    }
32
33    pub(crate) fn update_snapshot(&mut self, mut snapshot: MemoryResource) {
34        snapshot.normalize_capacity();
35        debug_assert_eq!(snapshot.validate(), Ok(()));
36        if self.snapshot.capacity_source != CapacitySource::User
37            || snapshot.capacity_source == CapacitySource::User
38        {
39            self.snapshot = snapshot;
40        }
41    }
42
43    pub(crate) fn apply_telemetry(
44        &mut self,
45        target: &MemoryResource,
46        snapshot: CapacitySnapshot,
47    ) -> bool {
48        if !self.snapshot.is_refreshable()
49            || self.snapshot.device_identity != target.device_identity
50        {
51            return false;
52        }
53        self.snapshot.apply_capacity_snapshot(snapshot);
54        debug_assert_eq!(self.snapshot.validate(), Ok(()));
55        true
56    }
57
58    fn try_reserve(&mut self, bytes: u64) -> MemoryResult<()> {
59        let physical_limit = self.snapshot.effective_available();
60        let Some(next) = self.reserved.checked_add(bytes) else {
61            return Err(self.budget_exceeded(bytes, physical_limit));
62        };
63        if next > physical_limit {
64            return Err(self.budget_exceeded(bytes, physical_limit));
65        }
66        self.reserved = next;
67        self.high_water = self.high_water.max(next);
68        Ok(())
69    }
70
71    fn release(&mut self, bytes: u64) {
72        self.reserved = self.reserved.saturating_sub(bytes);
73    }
74    fn budget_exceeded(&self, requested: u64, physical_limit: u64) -> MemoryError {
75        MemoryError::BudgetExceeded {
76            resource: self.snapshot.name.clone(),
77            requested,
78            remaining: physical_limit.saturating_sub(self.reserved),
79        }
80    }
81}
82
83/// A resolved, reservable limit within one physical resource.
84#[derive(Clone, Debug)]
85pub struct MemoryPool {
86    pub(crate) inner: Arc<MemoryPoolInner>,
87}
88
89#[derive(Debug)]
90pub(crate) struct MemoryPoolInner {
91    pub(crate) requested: MemoryBudget,
92    pub(crate) accounting: Arc<ReservationAccount>,
93}
94
95#[derive(Debug)]
96pub(crate) struct ReservationAccount {
97    pub(crate) state: Weak<MemoryStateInner>,
98    pub(crate) resource_id: String,
99    pub(crate) capacity: u64,
100    reserved: AtomicU64,
101    high_water: AtomicU64,
102}
103
104impl ReservationAccount {
105    pub(crate) fn new(state: Weak<MemoryStateInner>, resource_id: String, capacity: u64) -> Self {
106        Self {
107            state,
108            resource_id,
109            capacity,
110            reserved: AtomicU64::new(0),
111            high_water: AtomicU64::new(0),
112        }
113    }
114
115    fn try_reserve(self: &Arc<Self>, bytes: u64) -> MemoryResult<ReservationToken> {
116        let state = self.state.upgrade();
117        let mut resources = state
118            .as_ref()
119            .map(|state| state.resources.lock().unwrap_or_else(|e| e.into_inner()));
120        let next = self.try_reserve_local(bytes)?;
121        if let Some(ledger) = resources
122            .as_mut()
123            .and_then(|resources| resources.get_mut(&self.resource_id))
124            && let Err(error) = ledger.try_reserve(bytes)
125        {
126            self.release_local(bytes);
127            return Err(error);
128        }
129        update_max(&self.high_water, next);
130        Ok(ReservationToken {
131            account: Arc::clone(self),
132            bytes,
133        })
134    }
135
136    fn try_reserve_local(&self, bytes: u64) -> MemoryResult<u64> {
137        self.reserved
138            .try_update(Ordering::AcqRel, Ordering::Acquire, |current| {
139                current
140                    .checked_add(bytes)
141                    .filter(|&next| next <= self.capacity)
142            })
143            .map(|previous| previous + bytes)
144            .map_err(|_| self.budget_exceeded(bytes))
145    }
146
147    fn release(&self, bytes: u64) {
148        if let Some(state) = self.state.upgrade() {
149            let mut resources = state.resources.lock().unwrap_or_else(|e| e.into_inner());
150            if let Some(ledger) = resources.get_mut(&self.resource_id) {
151                ledger.release(bytes);
152            }
153        }
154        self.release_local(bytes);
155    }
156
157    fn release_local(&self, bytes: u64) {
158        self.reserved.fetch_sub(bytes, Ordering::AcqRel);
159    }
160    fn budget_exceeded(&self, requested: u64) -> MemoryError {
161        MemoryError::BudgetExceeded {
162            resource: self.resource_id.clone(),
163            requested,
164            remaining: self
165                .capacity
166                .saturating_sub(self.reserved.load(Ordering::Acquire)),
167        }
168    }
169}
170
171#[derive(Debug)]
172struct ReservationToken {
173    account: Arc<ReservationAccount>,
174    bytes: u64,
175}
176impl Drop for ReservationToken {
177    fn drop(&mut self) {
178        self.account.release(self.bytes);
179    }
180}
181
182impl MemoryPool {
183    /// Whether two pool handles charge the same reservation account.
184    pub fn shares_reservations_with(&self, other: &Self) -> bool {
185        Arc::ptr_eq(&self.inner.accounting, &other.inner.accounting)
186    }
187    /// Requested budget specification.
188    pub fn requested(&self) -> MemoryBudget {
189        self.inner.requested
190    }
191    /// Resolved pool capacity in bytes.
192    pub fn capacity(&self) -> u64 {
193        self.inner.accounting.capacity
194    }
195    /// Currently reserved bytes.
196    pub fn reserved(&self) -> u64 {
197        self.inner.accounting.reserved.load(Ordering::Acquire)
198    }
199    /// Remaining reservable bytes.
200    pub fn remaining(&self) -> u64 {
201        self.capacity().saturating_sub(self.reserved())
202    }
203    /// Highest concurrent reservation observed by this pool.
204    pub fn high_water(&self) -> u64 {
205        self.inner.accounting.high_water.load(Ordering::Acquire)
206    }
207    /// Attempts to reserve `bytes` until the returned lease is dropped.
208    ///
209    /// # Errors
210    ///
211    /// Returns [`MemoryError::BudgetExceeded`] if the local or shared resource
212    /// limit lacks sufficient remaining capacity.
213    pub fn reserve(&self, bytes: u64) -> MemoryResult<MemoryLease> {
214        Ok(MemoryLease {
215            inner: Arc::new(MemoryLeaseInner {
216                reservation: self.inner.accounting.try_reserve(bytes)?,
217            }),
218        })
219    }
220    /// Returns a snapshot of this pool's planning and usage.
221    pub fn report(&self) -> MemoryPoolReport {
222        MemoryPoolReport {
223            resource_id: self.inner.accounting.resource_id.clone(),
224            requested: self.requested(),
225            effective_bytes: self.capacity(),
226            reserved_bytes: self.reserved(),
227            remaining_bytes: self.remaining(),
228            high_water_bytes: self.high_water(),
229        }
230    }
231}
232
233/// A live memory reservation. Clones share one reservation.
234#[derive(Clone, Debug)]
235pub struct MemoryLease {
236    inner: Arc<MemoryLeaseInner>,
237}
238#[derive(Debug)]
239struct MemoryLeaseInner {
240    reservation: ReservationToken,
241}
242impl MemoryLease {
243    /// Reserved byte count.
244    pub fn bytes(&self) -> u64 {
245        self.inner.reservation.bytes
246    }
247}
248
249fn update_max(value: &AtomicU64, candidate: u64) {
250    let mut current = value.load(Ordering::Acquire);
251    while candidate > current {
252        match value.compare_exchange_weak(current, candidate, Ordering::AcqRel, Ordering::Acquire) {
253            Ok(_) => break,
254            Err(observed) => current = observed,
255        }
256    }
257}
258
259#[cfg(test)]
260mod tests {
261    use super::*;
262    use crate::{MemoryResourceKind, MemoryState};
263    use std::thread;
264
265    fn resource() -> MemoryResource {
266        MemoryResource {
267            id: "test".into(),
268            name: "Test".into(),
269            kind: MemoryResourceKind::Device,
270            total_bytes: Some(1_000),
271            available_bytes: Some(500),
272            capacity_source: CapacitySource::User,
273            device_identity: None,
274        }
275    }
276    fn resource_reserved(state: &MemoryState) -> u64 {
277        state
278            .report()
279            .resources
280            .into_iter()
281            .find(|r| r.resource.id == "test")
282            .unwrap()
283            .laddu_reserved_bytes
284    }
285
286    fn concurrent_reservations(pool: &MemoryPool) -> Vec<MemoryLease> {
287        (0..32)
288            .map(|_| {
289                let pool = pool.clone();
290                thread::spawn(move || pool.reserve(10))
291            })
292            .collect::<Vec<_>>()
293            .into_iter()
294            .filter_map(|handle| handle.join().unwrap().ok())
295            .collect()
296    }
297
298    #[test]
299    fn cloned_leases_release_once_and_preserve_high_water() {
300        let state = MemoryState::discover();
301        state.insert_device_snapshot(resource());
302        let pool = state.pool("test", MemoryBudget::Bytes(300)).unwrap();
303        let zero = pool.reserve(0).unwrap();
304        assert_eq!((pool.reserved(), resource_reserved(&state)), (0, 0));
305        drop(zero);
306        let lease = pool.reserve(120).unwrap();
307        let clone = lease.clone();
308        assert_eq!((pool.reserved(), resource_reserved(&state)), (120, 120));
309        drop(lease);
310        assert_eq!(pool.reserved(), 120);
311        drop(clone);
312        assert_eq!((pool.reserved(), pool.high_water()), (0, 120));
313        assert_eq!(resource_reserved(&state), 0);
314    }
315
316    #[test]
317    fn leases_enforce_release_and_align_shared_capacity() {
318        let state = MemoryState::discover();
319        state.insert_device_snapshot(resource());
320        let first = state.pool("test", MemoryBudget::Bytes(400)).unwrap();
321        let second = state.pool("test", MemoryBudget::Bytes(400)).unwrap();
322        let first_lease = first.reserve(300).unwrap();
323        assert!(second.reserve(201).is_err());
324        assert_eq!((second.reserved(), resource_reserved(&state)), (0, 300));
325        let second_lease = second.reserve(200).unwrap();
326        assert_eq!(resource_reserved(&state), 500);
327        drop(first_lease);
328        drop(second_lease);
329        assert_eq!(resource_reserved(&state), 0);
330    }
331
332    #[test]
333    fn overflow_failure_has_no_side_effects() {
334        let state = MemoryState::discover();
335        let maximum = MemoryResource {
336            total_bytes: Some(u64::MAX),
337            available_bytes: Some(u64::MAX),
338            ..resource()
339        };
340        state.insert_device_snapshot(maximum);
341        let pool = state.pool("test", MemoryBudget::Bytes(u64::MAX)).unwrap();
342        let lease = pool.reserve(u64::MAX).unwrap();
343        assert!(pool.reserve(1).is_err());
344        assert_eq!((pool.reserved(), pool.high_water()), (u64::MAX, u64::MAX));
345        drop(lease);
346        assert_eq!(pool.reserved(), 0);
347    }
348
349    #[test]
350    fn concurrent_reservations_are_atomic_with_and_without_state() {
351        let state = MemoryState::discover();
352        state.insert_device_snapshot(resource());
353        let attached = state.pool("test", MemoryBudget::Bytes(100)).unwrap();
354        let leases = concurrent_reservations(&attached);
355        assert_eq!((leases.len(), attached.reserved()), (10, 100));
356        drop(leases);
357        assert_eq!(attached.reserved(), 0);
358
359        let detached = state.pool("test", MemoryBudget::Bytes(100)).unwrap();
360        drop(state);
361        let leases = concurrent_reservations(&detached);
362        assert_eq!((leases.len(), detached.reserved()), (10, 100));
363        drop(leases);
364        assert_eq!(detached.reserved(), 0);
365    }
366}