reifydb_core/util/
budget.rs1use std::sync::atomic::{AtomicU64, Ordering};
5
6use reifydb_value::byte_size::ByteSize;
7
8pub struct MemoryBudget {
9 used: AtomicU64,
10 limit: AtomicU64,
11}
12
13impl MemoryBudget {
14 pub fn new(limit: ByteSize) -> Self {
15 Self {
16 used: AtomicU64::new(0),
17 limit: AtomicU64::new(limit.as_bytes()),
18 }
19 }
20
21 pub fn charge(&self, bytes: ByteSize) {
22 self.used.fetch_add(bytes.as_bytes(), Ordering::Relaxed);
23 }
24
25 pub fn try_charge(&self, bytes: ByteSize) -> bool {
26 let amount = bytes.as_bytes();
27 let limit = self.limit.load(Ordering::Relaxed);
28 let mut current = self.used.load(Ordering::Relaxed);
29 loop {
30 let next = current.saturating_add(amount);
31 if next > limit {
32 return false;
33 }
34 match self.used.compare_exchange_weak(current, next, Ordering::Relaxed, Ordering::Relaxed) {
35 Ok(_) => return true,
36 Err(observed) => current = observed,
37 }
38 }
39 }
40
41 pub fn release(&self, bytes: ByteSize) {
42 let amount = bytes.as_bytes();
43 let mut current = self.used.load(Ordering::Relaxed);
44 loop {
45 let next = current.saturating_sub(amount);
46 match self.used.compare_exchange_weak(current, next, Ordering::Relaxed, Ordering::Relaxed) {
47 Ok(_) => return,
48 Err(observed) => current = observed,
49 }
50 }
51 }
52
53 pub fn over_budget(&self) -> bool {
54 self.used.load(Ordering::Relaxed) > self.limit.load(Ordering::Relaxed)
55 }
56
57 pub fn used(&self) -> ByteSize {
58 ByteSize::from_bytes(self.used.load(Ordering::Relaxed))
59 }
60
61 pub fn limit(&self) -> ByteSize {
62 ByteSize::from_bytes(self.limit.load(Ordering::Relaxed))
63 }
64
65 pub fn reset(&self) {
66 self.used.store(0, Ordering::Relaxed);
67 }
68}
69
70#[cfg(test)]
71mod tests {
72 use std::{sync::Arc, thread};
73
74 use reifydb_value::byte_size::ByteSize;
75
76 use super::MemoryBudget;
77
78 #[test]
79 fn charge_and_release_track_used() {
80 let budget = MemoryBudget::new(ByteSize::from_kib(4));
81 budget.charge(ByteSize::from_kib(1));
82 budget.charge(ByteSize::from_kib(2));
83 assert_eq!(budget.used(), ByteSize::from_kib(3));
84 budget.release(ByteSize::from_kib(2));
85 assert_eq!(budget.used(), ByteSize::from_kib(1));
86 }
87
88 #[test]
89 fn over_budget_trips_only_above_limit() {
90 let budget = MemoryBudget::new(ByteSize::from_kib(2));
91 budget.charge(ByteSize::from_kib(2));
92 assert!(!budget.over_budget(), "used == limit is within budget");
93 budget.charge(ByteSize::from_bytes(1));
94 assert!(budget.over_budget(), "one byte over the limit trips it");
95 }
96
97 #[test]
98 fn release_saturates_at_zero() {
99 let budget = MemoryBudget::new(ByteSize::from_kib(4));
100 budget.charge(ByteSize::from_kib(1));
101 budget.release(ByteSize::from_kib(5));
102 assert_eq!(budget.used(), ByteSize::ZERO, "release never underflows the counter");
103 }
104
105 #[test]
106 fn try_charge_commits_only_when_it_fits() {
107 let budget = MemoryBudget::new(ByteSize::from_kib(4));
108 assert!(budget.try_charge(ByteSize::from_kib(3)), "charge within limit succeeds");
109 assert_eq!(budget.used(), ByteSize::from_kib(3));
110 assert!(budget.try_charge(ByteSize::from_kib(1)), "charge up to exactly the limit succeeds");
111 assert_eq!(budget.used(), ByteSize::from_kib(4));
112 }
113
114 #[test]
115 fn try_charge_rejects_and_leaves_used_unchanged() {
116 let budget = MemoryBudget::new(ByteSize::from_kib(4));
117 budget.charge(ByteSize::from_kib(3));
118 assert!(!budget.try_charge(ByteSize::from_kib(2)), "a charge that would exceed the limit is rejected");
119 assert_eq!(budget.used(), ByteSize::from_kib(3), "a rejected charge must not mutate used");
120 }
121
122 #[test]
123 fn reset_zeroes_used_and_keeps_limit() {
124 let budget = MemoryBudget::new(ByteSize::from_kib(4));
125 budget.charge(ByteSize::from_kib(3));
126 budget.reset();
127 assert_eq!(budget.used(), ByteSize::ZERO);
128 assert_eq!(budget.limit(), ByteSize::from_kib(4));
129 }
130
131 #[test]
132 fn concurrent_try_charge_never_exceeds_limit() {
133 let budget = Arc::new(MemoryBudget::new(ByteSize::from_bytes(10_000)));
134 let threads = 16;
135 let attempts_per_thread = 2_000;
136 let charge_amount = 7u64;
137
138 let handles: Vec<_> = (0..threads)
139 .map(|_| {
140 let budget = budget.clone();
141 thread::spawn(move || {
142 (0..attempts_per_thread)
143 .filter(|_| budget.try_charge(ByteSize::from_bytes(charge_amount)))
144 .count()
145 })
146 })
147 .collect();
148
149 let total_successes: u64 = handles.into_iter().map(|h| h.join().unwrap() as u64).sum();
150
151 assert!(
152 budget.used().as_bytes() <= budget.limit().as_bytes(),
153 "the CAS retry loop must never let concurrent charges overshoot the limit"
154 );
155 assert_eq!(
156 budget.used().as_bytes(),
157 total_successes * charge_amount,
158 "used must equal exactly the sum of successful charges, with no lost or duplicated updates"
159 );
160 }
161}