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    /// Requested budget specification.
184    pub fn requested(&self) -> MemoryBudget {
185        self.inner.requested
186    }
187    /// Resolved pool capacity in bytes.
188    pub fn capacity(&self) -> u64 {
189        self.inner.accounting.capacity
190    }
191    /// Currently reserved bytes.
192    pub fn reserved(&self) -> u64 {
193        self.inner.accounting.reserved.load(Ordering::Acquire)
194    }
195    /// Remaining reservable bytes.
196    pub fn remaining(&self) -> u64 {
197        self.capacity().saturating_sub(self.reserved())
198    }
199    /// Highest concurrent reservation observed by this pool.
200    pub fn high_water(&self) -> u64 {
201        self.inner.accounting.high_water.load(Ordering::Acquire)
202    }
203    /// Attempts to reserve `bytes` until the returned lease is dropped.
204    ///
205    /// # Errors
206    ///
207    /// Returns [`MemoryError::BudgetExceeded`] if the local or shared resource
208    /// limit lacks sufficient remaining capacity.
209    pub fn reserve(&self, bytes: u64) -> MemoryResult<MemoryLease> {
210        Ok(MemoryLease {
211            inner: Arc::new(MemoryLeaseInner {
212                reservation: self.inner.accounting.try_reserve(bytes)?,
213            }),
214        })
215    }
216    /// Returns a snapshot of this pool's planning and usage.
217    pub fn report(&self) -> MemoryPoolReport {
218        MemoryPoolReport {
219            resource_id: self.inner.accounting.resource_id.clone(),
220            requested: self.requested(),
221            effective_bytes: self.capacity(),
222            reserved_bytes: self.reserved(),
223            remaining_bytes: self.remaining(),
224            high_water_bytes: self.high_water(),
225        }
226    }
227}
228
229/// A live memory reservation. Clones share one reservation.
230#[derive(Clone, Debug)]
231pub struct MemoryLease {
232    inner: Arc<MemoryLeaseInner>,
233}
234#[derive(Debug)]
235struct MemoryLeaseInner {
236    reservation: ReservationToken,
237}
238impl MemoryLease {
239    /// Reserved byte count.
240    pub fn bytes(&self) -> u64 {
241        self.inner.reservation.bytes
242    }
243}
244
245fn update_max(value: &AtomicU64, candidate: u64) {
246    let mut current = value.load(Ordering::Acquire);
247    while candidate > current {
248        match value.compare_exchange_weak(current, candidate, Ordering::AcqRel, Ordering::Acquire) {
249            Ok(_) => break,
250            Err(observed) => current = observed,
251        }
252    }
253}
254
255#[cfg(test)]
256mod tests {
257    use super::*;
258    use crate::{MemoryResourceKind, MemoryState};
259    use std::thread;
260
261    fn resource() -> MemoryResource {
262        MemoryResource {
263            id: "test".into(),
264            name: "Test".into(),
265            kind: MemoryResourceKind::Device,
266            total_bytes: Some(1_000),
267            available_bytes: Some(500),
268            capacity_source: CapacitySource::User,
269            device_identity: None,
270        }
271    }
272    fn resource_reserved(state: &MemoryState) -> u64 {
273        state
274            .report()
275            .resources
276            .into_iter()
277            .find(|r| r.resource.id == "test")
278            .unwrap()
279            .laddu_reserved_bytes
280    }
281
282    fn concurrent_reservations(pool: &MemoryPool) -> Vec<MemoryLease> {
283        (0..32)
284            .map(|_| {
285                let pool = pool.clone();
286                thread::spawn(move || pool.reserve(10))
287            })
288            .collect::<Vec<_>>()
289            .into_iter()
290            .filter_map(|handle| handle.join().unwrap().ok())
291            .collect()
292    }
293
294    #[test]
295    fn cloned_leases_release_once_and_preserve_high_water() {
296        let state = MemoryState::discover();
297        state.insert_device_snapshot(resource());
298        let pool = state.pool("test", MemoryBudget::Bytes(300)).unwrap();
299        let zero = pool.reserve(0).unwrap();
300        assert_eq!((pool.reserved(), resource_reserved(&state)), (0, 0));
301        drop(zero);
302        let lease = pool.reserve(120).unwrap();
303        let clone = lease.clone();
304        assert_eq!((pool.reserved(), resource_reserved(&state)), (120, 120));
305        drop(lease);
306        assert_eq!(pool.reserved(), 120);
307        drop(clone);
308        assert_eq!((pool.reserved(), pool.high_water()), (0, 120));
309        assert_eq!(resource_reserved(&state), 0);
310    }
311
312    #[test]
313    fn leases_enforce_release_and_align_shared_capacity() {
314        let state = MemoryState::discover();
315        state.insert_device_snapshot(resource());
316        let first = state.pool("test", MemoryBudget::Bytes(400)).unwrap();
317        let second = state.pool("test", MemoryBudget::Bytes(400)).unwrap();
318        let first_lease = first.reserve(300).unwrap();
319        assert!(second.reserve(201).is_err());
320        assert_eq!((second.reserved(), resource_reserved(&state)), (0, 300));
321        let second_lease = second.reserve(200).unwrap();
322        assert_eq!(resource_reserved(&state), 500);
323        drop(first_lease);
324        drop(second_lease);
325        assert_eq!(resource_reserved(&state), 0);
326    }
327
328    #[test]
329    fn overflow_failure_has_no_side_effects() {
330        let state = MemoryState::discover();
331        let maximum = MemoryResource {
332            total_bytes: Some(u64::MAX),
333            available_bytes: Some(u64::MAX),
334            ..resource()
335        };
336        state.insert_device_snapshot(maximum);
337        let pool = state.pool("test", MemoryBudget::Bytes(u64::MAX)).unwrap();
338        let lease = pool.reserve(u64::MAX).unwrap();
339        assert!(pool.reserve(1).is_err());
340        assert_eq!((pool.reserved(), pool.high_water()), (u64::MAX, u64::MAX));
341        drop(lease);
342        assert_eq!(pool.reserved(), 0);
343    }
344
345    #[test]
346    fn concurrent_reservations_are_atomic_with_and_without_state() {
347        let state = MemoryState::discover();
348        state.insert_device_snapshot(resource());
349        let attached = state.pool("test", MemoryBudget::Bytes(100)).unwrap();
350        let leases = concurrent_reservations(&attached);
351        assert_eq!((leases.len(), attached.reserved()), (10, 100));
352        drop(leases);
353        assert_eq!(attached.reserved(), 0);
354
355        let detached = state.pool("test", MemoryBudget::Bytes(100)).unwrap();
356        drop(state);
357        let leases = concurrent_reservations(&detached);
358        assert_eq!((leases.len(), detached.reserved()), (10, 100));
359        drop(leases);
360        assert_eq!(detached.reserved(), 0);
361    }
362}