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#[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 pub fn requested(&self) -> MemoryBudget {
185 self.inner.requested
186 }
187 pub fn capacity(&self) -> u64 {
189 self.inner.accounting.capacity
190 }
191 pub fn reserved(&self) -> u64 {
193 self.inner.accounting.reserved.load(Ordering::Acquire)
194 }
195 pub fn remaining(&self) -> u64 {
197 self.capacity().saturating_sub(self.reserved())
198 }
199 pub fn high_water(&self) -> u64 {
201 self.inner.accounting.high_water.load(Ordering::Acquire)
202 }
203 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 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#[derive(Clone, Debug)]
231pub struct MemoryLease {
232 inner: Arc<MemoryLeaseInner>,
233}
234#[derive(Debug)]
235struct MemoryLeaseInner {
236 reservation: ReservationToken,
237}
238impl MemoryLease {
239 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}