1use dashmap::DashMap;
4use parking_lot::Mutex;
5use std::sync::atomic::{AtomicI64, Ordering};
6use std::sync::Arc;
7
8#[derive(Debug, thiserror::Error, PartialEq, Eq)]
10#[non_exhaustive]
11pub enum StateError {
12 #[error("counter overflow for key {0:?}")]
14 Overflow(String),
15 #[error("state backend unavailable: {0}")]
18 Backend(String),
19}
20
21#[derive(Debug, Clone, PartialEq, Eq)]
23pub struct Spend {
24 pub key: String,
26 pub amount: u64,
28 pub limit: u64,
30}
31
32pub trait StateStore: Send + Sync {
36 fn add(&self, key: &str, delta: u64) -> Result<u64, StateError>;
38
39 fn get(&self, key: &str) -> Result<u64, StateError>;
41
42 fn try_spend(&self, key: &str, amount: u64, limit: u64) -> Result<bool, StateError>;
46
47 fn try_spend_many(&self, spends: &[Spend]) -> Result<Option<usize>, StateError>;
56
57 fn remove(&self, key: &str);
59
60 fn refund(&self, key: &str, amount: u64) {
77 let _ = (key, amount);
78 }
79
80 fn remove_prefix(&self, prefix: &str) {
85 let _ = prefix;
86 }
87}
88
89#[derive(Debug, Default)]
92pub struct InMemoryStore {
93 counters: DashMap<String, Arc<AtomicI64>>,
94 transaction_lock: Mutex<()>,
95}
96
97impl InMemoryStore {
98 pub fn new() -> Self {
100 Self::default()
101 }
102
103 fn cell(&self, key: &str) -> Arc<AtomicI64> {
104 self.counters
105 .entry(key.to_owned())
106 .or_insert_with(|| Arc::new(AtomicI64::new(0)))
107 .clone()
108 }
109}
110
111pub(crate) const COUNTER_MAX: i64 = av_core::error::JCS_SAFE_MAX as i64;
120
121impl StateStore for InMemoryStore {
122 fn add(&self, key: &str, delta: u64) -> Result<u64, StateError> {
123 let _transaction = self.transaction_lock.lock();
124 let delta = i64::try_from(delta).map_err(|_| StateError::Overflow(key.to_owned()))?;
125 let cell = self.cell(key);
126 let prev = cell.load(Ordering::Acquire);
127 let new = prev
128 .checked_add(delta)
129 .filter(|v| *v <= COUNTER_MAX)
130 .ok_or_else(|| StateError::Overflow(key.to_owned()))?;
131 cell.store(new, Ordering::Release);
133 Ok(u64::try_from(new).unwrap_or(0))
134 }
135
136 fn get(&self, key: &str) -> Result<u64, StateError> {
137 match self.counters.get(key) {
138 None => Ok(0),
139 Some(cell) => {
140 let raw = cell.load(Ordering::Acquire);
141 if raw < 0 {
142 return Err(StateError::Overflow(key.to_owned()));
143 }
144 u64::try_from(raw).map_err(|_| StateError::Overflow(key.to_owned()))
145 }
146 }
147 }
148
149 fn try_spend(&self, key: &str, amount: u64, limit: u64) -> Result<bool, StateError> {
150 Ok(self
151 .try_spend_many(&[Spend {
152 key: key.to_owned(),
153 amount,
154 limit,
155 }])?
156 .is_none())
157 }
158
159 fn try_spend_many(&self, spends: &[Spend]) -> Result<Option<usize>, StateError> {
160 let _transaction = self.transaction_lock.lock();
161 let mut seen = std::collections::HashSet::with_capacity(spends.len());
166 for spend in spends {
167 if !seen.insert(spend.key.as_str()) {
168 return Err(StateError::Backend(format!(
169 "try_spend_many received duplicate key {:?}",
170 spend.key,
171 )));
172 }
173 }
174 let mut prepared = Vec::with_capacity(spends.len());
175 for (index, spend) in spends.iter().enumerate() {
176 if spend.amount > av_core::error::JCS_SAFE_MAX {
185 return Err(StateError::Overflow(spend.key.clone()));
186 }
187 if spend.limit > av_core::error::JCS_SAFE_MAX {
188 return Err(StateError::Overflow(spend.key.clone()));
189 }
190 let amount = i64::try_from(spend.amount).map_err(|_| StateError::Overflow(spend.key.clone()))?;
191 let limit = i64::try_from(spend.limit).map_err(|_| StateError::Overflow(spend.key.clone()))?;
192 let cell = self.cell(&spend.key);
193 let current = cell.load(Ordering::Acquire);
194 let next = current
195 .checked_add(amount)
196 .ok_or_else(|| StateError::Overflow(spend.key.clone()))?;
197 if next > limit {
198 return Ok(Some(index));
199 }
200 prepared.push((cell, amount));
201 }
202 for (cell, amount) in prepared {
203 cell.fetch_add(amount, Ordering::AcqRel);
204 }
205 Ok(None)
206 }
207
208 fn remove(&self, key: &str) {
209 let _transaction = self.transaction_lock.lock();
210 self.counters.remove(key);
211 }
212
213 fn refund(&self, key: &str, amount: u64) {
232 let _transaction = self.transaction_lock.lock();
233 let Some(cell) = self.counters.get(key).map(|entry| Arc::clone(entry.value())) else {
234 return;
235 };
236 let prev = cell.load(Ordering::Acquire);
237 let amount = i64::try_from(amount).unwrap_or(i64::MAX);
238 let next = prev.saturating_sub(amount).max(0);
239 cell.store(next, Ordering::Release);
240 }
241
242 fn remove_prefix(&self, prefix: &str) {
243 let _transaction = self.transaction_lock.lock();
244 self.counters.retain(|key, _| !key.starts_with(prefix));
245 }
246}
247
248#[cfg(test)]
249mod tests {
250 #![allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
251
252 use super::*;
253
254 #[test]
255 fn remove_prefix_drops_only_matching_keys() {
256 let store = InMemoryStore::new();
257 store.add("budget:{aaaa}:tokens", 5).unwrap();
258 store.add("budget:{aaaa}:tool:db_write", 1).unwrap();
259 store.add("budget:{bbbb}:tokens", 7).unwrap();
260 store.remove_prefix("budget:{aaaa}:");
261 assert_eq!(store.get("budget:{aaaa}:tokens").unwrap(), 0);
262 assert_eq!(store.get("budget:{aaaa}:tool:db_write").unwrap(), 0);
263 assert_eq!(
264 store.get("budget:{bbbb}:tokens").unwrap(),
265 7,
266 "other sessions' counters must survive a prefix removal",
267 );
268 assert_eq!(store.counters.len(), 1, "removed cells must actually be freed");
269 }
270
271 #[test]
272 fn add_overflow_rollback_is_never_visible_to_concurrent_get() {
273 use std::sync::{Arc, Barrier};
279 use std::thread;
280
281 let store = Arc::new(InMemoryStore::new());
282 store.add("k", 5).unwrap();
283 let barrier = Arc::new(Barrier::new(3));
284
285 let s1 = Arc::clone(&store);
287 let b1 = Arc::clone(&barrier);
288 let writer = thread::spawn(move || {
289 b1.wait();
290 for _ in 0..500 {
291 let _ = s1.add("k", u64::MAX / 2); }
293 });
294
295 let s2 = Arc::clone(&store);
297 let b2 = Arc::clone(&barrier);
298 let reader1 = thread::spawn(move || {
299 b2.wait();
300 for _ in 0..2_000 {
301 let v = s2.get("k");
302 assert!(v.is_ok(), "spurious Overflow from concurrent get: {v:?}");
303 }
304 });
305
306 let s3 = Arc::clone(&store);
307 let b3 = Arc::clone(&barrier);
308 let reader2 = thread::spawn(move || {
309 b3.wait();
310 for _ in 0..2_000 {
311 let v = s3.get("k");
312 assert!(v.is_ok(), "spurious Overflow from concurrent get: {v:?}");
313 }
314 });
315
316 writer.join().unwrap();
317 reader1.join().unwrap();
318 reader2.join().unwrap();
319 }
320
321 #[test]
322 fn poisoned_negative_value_surfaces_as_overflow_not_silent_zero() {
323 let s = InMemoryStore::new();
328 let cell = s.cell("k");
329 cell.store(-1, Ordering::Release);
330 match s.get("k") {
331 Err(StateError::Overflow(_)) => (),
332 other => panic!("expected Overflow on negative counter, got {other:?}"),
333 }
334 }
335
336 #[test]
337 fn add_and_get() {
338 let s = InMemoryStore::new();
339 assert_eq!(s.get("k").unwrap(), 0);
340 assert_eq!(s.add("k", 5).unwrap(), 5);
341 assert_eq!(s.add("k", 3).unwrap(), 8);
342 assert_eq!(s.get("k").unwrap(), 8);
343 s.remove("k");
344 assert_eq!(s.get("k").unwrap(), 0);
345 }
346
347 #[test]
348 fn try_spend_respects_limit_exactly() {
349 let s = InMemoryStore::new();
350 assert!(s.try_spend("b", 3, 3).unwrap()); assert!(!s.try_spend("b", 1, 3).unwrap()); assert_eq!(s.get("b").unwrap(), 3, "refused spend must not record");
353 }
354
355 #[test]
356 fn zero_amount_spend_is_free() {
357 let s = InMemoryStore::new();
358 assert!(s.try_spend("z", 0, 0).unwrap());
359 assert_eq!(s.get("z").unwrap(), 0);
360 }
361
362 #[test]
363 fn overflow_is_loud_not_wrapping() {
364 let s = InMemoryStore::new();
365 assert!(matches!(s.add("o", u64::MAX), Err(StateError::Overflow(_))));
366 }
367
368 #[test]
371 fn concurrent_spend_never_exceeds_budget() {
372 let s = Arc::new(InMemoryStore::new());
373 let limit = 10_000u64;
374 let mut handles = Vec::new();
375 for _ in 0..64 {
376 let s = Arc::clone(&s);
377 handles.push(std::thread::spawn(move || {
378 let mut granted = 0u64;
379 for _ in 0..1000 {
380 if s.try_spend("shared", 1, limit).unwrap() {
381 granted += 1;
382 }
383 }
384 granted
385 }));
386 }
387 let total: u64 = handles.into_iter().map(|h| h.join().unwrap()).sum();
388 assert_eq!(total, limit, "grants must equal the budget exactly");
389 assert_eq!(s.get("shared").unwrap(), limit);
390 }
391
392 #[test]
394 fn concurrent_mixed_amounts_never_over_cap() {
395 let s = Arc::new(InMemoryStore::new());
396 let limit = 5_000u64;
397 let mut handles = Vec::new();
398 for t in 0..32 {
399 let s = Arc::clone(&s);
400 handles.push(std::thread::spawn(move || {
401 let mut spent = 0u64;
402 let amount = (t % 7) + 1;
403 for _ in 0..500 {
404 if s.try_spend("cap", amount, limit).unwrap() {
405 spent += amount;
406 }
407 }
408 spent
409 }));
410 }
411 let total: u64 = handles.into_iter().map(|h| h.join().unwrap()).sum();
412 assert!(total <= limit, "over-spend: {total} > {limit}");
413 assert_eq!(s.get("cap").unwrap(), total);
414 }
415
416 #[test]
423 fn try_spend_many_refuses_duplicate_keys() {
424 let s = InMemoryStore::new();
425 let outcome = s.try_spend_many(&[
426 Spend {
427 key: "budget".to_owned(),
428 amount: 60,
429 limit: 100,
430 },
431 Spend {
432 key: "budget".to_owned(),
433 amount: 60,
434 limit: 100,
435 },
436 ]);
437 match outcome {
438 Err(StateError::Backend(reason)) => {
439 assert!(reason.contains("duplicate key"), "wrong reason: {reason}");
440 }
441 other => panic!("must reject duplicate keys, got {other:?}"),
442 }
443 assert_eq!(
444 s.get("budget").unwrap(),
445 0,
446 "no partial spend must have been committed",
447 );
448 }
449
450 #[test]
453 fn try_spend_many_distinct_keys_still_commits_atomically() {
454 let s = InMemoryStore::new();
455 assert_eq!(
456 s.try_spend_many(&[
457 Spend {
458 key: "a".to_owned(),
459 amount: 3,
460 limit: 10,
461 },
462 Spend {
463 key: "b".to_owned(),
464 amount: 4,
465 limit: 10,
466 },
467 ])
468 .unwrap(),
469 None,
470 );
471 assert_eq!(s.get("a").unwrap(), 3);
472 assert_eq!(s.get("b").unwrap(), 4);
473 }
474
475 #[test]
483 fn try_spend_many_rejects_limits_past_counter_max() {
484 let s = InMemoryStore::new();
485 let outcome = s.try_spend_many(&[Spend {
486 key: "a".to_owned(),
487 amount: 1,
488 limit: av_core::error::JCS_SAFE_MAX + 1,
489 }]);
490 assert!(
491 matches!(outcome, Err(StateError::Overflow(_))),
492 "expected Overflow rejection, got {outcome:?}"
493 );
494 }
495
496 #[test]
497 fn try_spend_many_rejects_amounts_past_counter_max() {
498 let s = InMemoryStore::new();
499 let outcome = s.try_spend_many(&[Spend {
500 key: "a".to_owned(),
501 amount: av_core::error::JCS_SAFE_MAX + 1,
502 limit: u64::MAX,
503 }]);
504 assert!(
505 matches!(outcome, Err(StateError::Overflow(_))),
506 "expected Overflow rejection, got {outcome:?}"
507 );
508 }
509
510 #[test]
520 fn refund_after_remove_prefix_does_not_resurrect_cells() {
521 let s = InMemoryStore::new();
522 s.add("budget:{aaaa}:tool:db_write", 1).unwrap();
523 s.add("budget:{aaaa}:total_calls", 1).unwrap();
524 s.add("budget:{aaaa}:payout", 500_000).unwrap();
525 s.remove_prefix("budget:{aaaa}:");
528 assert_eq!(s.counters.len(), 0, "prefix sweep must have cleared all");
529 s.refund("budget:{aaaa}:tool:db_write", 1);
533 s.refund("budget:{aaaa}:total_calls", 1);
534 s.refund("budget:{aaaa}:payout", 500_000);
535 assert_eq!(
536 s.counters.len(),
537 0,
538 "refund on a swept session must not resurrect counter cells (attacker-choosable growth)"
539 );
540 }
541
542 #[test]
546 fn refund_on_live_session_still_compensates_exactly() {
547 let s = InMemoryStore::new();
548 s.add("budget:{live}:tool:db_write", 3).unwrap();
549 s.refund("budget:{live}:tool:db_write", 1);
550 assert_eq!(s.get("budget:{live}:tool:db_write").unwrap(), 2);
551 s.refund("budget:{live}:tool:db_write", 10);
553 assert_eq!(s.get("budget:{live}:tool:db_write").unwrap(), 0);
554 }
555}