1use std::collections::HashMap;
12use std::sync::atomic::{AtomicU64, Ordering};
13use std::sync::{Arc, LazyLock, Mutex, Weak};
14
15use crate::{Error, Result};
16
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
19pub struct Options {
20 pub increment: i64,
22 pub min: i64,
24 pub max: i64,
26 pub start: i64,
28 pub cycle: bool,
30}
31
32#[derive(Debug, Clone, Copy)]
34struct State {
35 counter: i64,
37 last: Option<i64>,
39 uses: u64,
41}
42
43#[derive(Debug)]
45pub struct Counter {
46 id: u64,
48 name: String,
50 options: Options,
52 state: Mutex<State>,
54}
55
56static COUNTERS: LazyLock<Mutex<HashMap<u64, Weak<Counter>>>> =
58 LazyLock::new(|| Mutex::new(HashMap::new()));
59
60static NEXT: AtomicU64 = AtomicU64::new(1);
62
63impl Counter {
64 #[must_use]
66 pub fn register(name: &str, options: Options) -> Arc<Self> {
67 let id = NEXT.fetch_add(1, Ordering::Relaxed);
68 let counter = Arc::new(Self {
69 id,
70 name: name.to_owned(),
71 options,
72 state: Mutex::new(State { counter: options.start, last: None, uses: 0 }),
73 });
74 let mut all = COUNTERS.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
75 all.retain(|_, counter| counter.strong_count() > 0);
76 all.insert(id, Arc::downgrade(&counter));
77 counter
78 }
79
80 #[must_use]
82 pub fn id(&self) -> u64 {
83 self.id
84 }
85
86 #[must_use]
88 pub fn options(&self) -> Options {
89 self.options
90 }
91
92 #[must_use]
95 pub fn counter(&self) -> i64 {
96 self.lock().counter
97 }
98
99 #[must_use]
101 pub fn last(&self) -> Option<i64> {
102 self.lock().last
103 }
104
105 #[must_use]
107 pub fn uses(&self) -> u64 {
108 self.lock().uses
109 }
110
111 fn lock(&self) -> std::sync::MutexGuard<'_, State> {
112 self.state.lock().unwrap_or_else(std::sync::PoisonError::into_inner)
113 }
114
115 pub fn next(&self) -> Result<i64> {
122 let mut state = self.lock();
123 self.advance(&mut state)
124 }
125
126 fn advance(&self, state: &mut State) -> Result<i64> {
127 let Options { increment, min, max, cycle, .. } = self.options;
128 let result = state.counter;
129 let (moved, overflow) = state.counter.overflowing_add(increment);
130 if !overflow {
131 state.counter = moved;
132 }
133 if cycle {
134 if overflow {
135 state.counter = if increment < 0 { max } else { min };
136 } else if state.counter < min {
137 state.counter = max;
138 } else if state.counter > max {
139 state.counter = min;
140 }
141 } else {
142 if result < min || (overflow && increment < 0) {
143 return Err(Error::sequence(format!(
144 "nextval: reached minimum value of sequence \"{}\" ({min})",
145 self.name
146 )));
147 }
148 if result > max || overflow {
149 return Err(Error::sequence(format!(
150 "nextval: reached maximum value of sequence \"\"{}\"\" ({max})",
151 self.name
152 )));
153 }
154 }
155 state.last = Some(result);
156 state.uses += 1;
157 Ok(result)
158 }
159
160 pub fn current(&self) -> Result<i64> {
166 self.lock()
167 .last
168 .ok_or_else(|| Error::sequence("currval: sequence is not yet defined in this session"))
169 }
170
171 pub fn set(&self, value: i64, called: bool) -> Result<i64> {
178 let Options { min, max, .. } = self.options;
179 if value < min || value > max {
180 return Err(Error::sequence(format!(
181 "setval: value {value} is out of bounds for sequence \"{}\" ({min}..{max})",
182 self.name
183 )));
184 }
185 let mut state = self.lock();
186 state.counter = value;
187 if !called {
188 state.uses += 1;
189 return Ok(value);
190 }
191 self.advance(&mut state)
192 }
193}
194
195pub fn lookup(id: u64) -> Result<Arc<Counter>> {
201 let all = COUNTERS.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
202 all.get(&id)
203 .and_then(Weak::upgrade)
204 .ok_or_else(|| Error::catalog("The sequence this statement uses no longer exists"))
205}
206
207#[cfg(test)]
208mod tests {
209 use super::*;
210
211 fn options(increment: i64, min: i64, max: i64, start: i64, cycle: bool) -> Options {
212 Options { increment, min, max, start, cycle }
213 }
214
215 #[test]
216 fn a_counter_counts_and_stops_at_its_maximum() {
217 let seq = Counter::register("seq", options(1, 1, 2, 1, false));
218 assert_eq!(seq.next().unwrap(), 1);
219 assert_eq!(seq.next().unwrap(), 2);
220 let error = seq.next().unwrap_err().to_string();
221 assert!(error.contains("reached maximum value of sequence \"\"seq\"\" (2)"), "{error}");
222 assert_eq!(seq.current().unwrap(), 2);
223 assert!(Arc::ptr_eq(&lookup(seq.id()).unwrap(), &seq));
224 }
225
226 #[test]
227 fn a_cycle_starts_again_from_the_other_end() {
228 let seq = Counter::register("seq", options(-1, 1, 3, 2, true));
229 let got: Vec<i64> = (0..4).map(|_| seq.next().unwrap()).collect();
230 assert_eq!(got, [2, 1, 3, 2]);
231 }
232
233 #[test]
234 fn setval_hands_out_the_value_when_called() {
235 let seq = Counter::register("seq", options(1, 1, i64::MAX, 1, false));
236 assert!(seq.current().is_err());
237 assert_eq!(seq.set(10, true).unwrap(), 10);
238 assert_eq!(seq.next().unwrap(), 11);
239 assert_eq!(seq.set(5, false).unwrap(), 5);
240 assert_eq!(seq.current().unwrap(), 11);
241 assert_eq!(seq.next().unwrap(), 5);
242 assert!(seq.set(0, true).is_err());
243 }
244}