Skip to main content

rudb_common/
sequence.rs

1//! The counter behind a `CREATE SEQUENCE`, and the registry the kernels find one in by number.
2//!
3//! The pin keeps a sequence's state on its catalog entry and changes it in place, so a value handed
4//! out by `nextval` stays handed out when the transaction that asked for it rolls back, and
5//! `currval` reads the last value the entry gave to anyone rather than one kept per session. The
6//! counter here is shared the same way. The catalog holds it behind an [`Arc`], the copy of the
7//! catalog a transaction keeps to roll back to holds the same one, and the binder turns
8//! `nextval('seq')` into a call that carries the counter's number rather than its name, which is
9//! what lets a kernel that knows nothing about catalogs reach it.
10
11use std::collections::HashMap;
12use std::sync::atomic::{AtomicU64, Ordering};
13use std::sync::{Arc, LazyLock, Mutex, Weak};
14
15use crate::{Error, Result};
16
17/// What a `CREATE SEQUENCE` settled, with every default already filled in.
18#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
19pub struct Options {
20    /// Added to the counter by every `nextval`, and never zero.
21    pub increment: i64,
22    /// The smallest value the sequence hands out.
23    pub min: i64,
24    /// The largest value the sequence hands out.
25    pub max: i64,
26    /// The first value handed out.
27    pub start: i64,
28    /// Whether running past one end starts again from the other rather than failing.
29    pub cycle: bool,
30}
31
32/// The part of a counter that changes.
33#[derive(Debug, Clone, Copy)]
34struct State {
35    /// The value the next `nextval` hands out.
36    counter: i64,
37    /// The value the last `nextval` or `setval` handed out.
38    last: Option<i64>,
39    /// How many values have been handed out.
40    uses: u64,
41}
42
43/// One sequence's counter.
44#[derive(Debug)]
45pub struct Counter {
46    /// The number the registry knows it by.
47    id: u64,
48    /// The sequence's name, for the messages.
49    name: String,
50    /// What it was created with.
51    options: Options,
52    /// Where it has got to.
53    state: Mutex<State>,
54}
55
56/// Every counter alive, by number.
57static COUNTERS: LazyLock<Mutex<HashMap<u64, Weak<Counter>>>> =
58    LazyLock::new(|| Mutex::new(HashMap::new()));
59
60/// The number the next counter gets.
61static NEXT: AtomicU64 = AtomicU64::new(1);
62
63impl Counter {
64    /// A new counter at its start value, registered so [`lookup`] finds it.
65    #[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    /// The number [`lookup`] finds this counter by.
81    #[must_use]
82    pub fn id(&self) -> u64 {
83        self.id
84    }
85
86    /// What it was created with.
87    #[must_use]
88    pub fn options(&self) -> Options {
89        self.options
90    }
91
92    /// The value the next `nextval` hands out, which is what the pin writes as `START` when it
93    /// prints the sequence back.
94    #[must_use]
95    pub fn counter(&self) -> i64 {
96        self.lock().counter
97    }
98
99    /// The value last handed out, if any has been.
100    #[must_use]
101    pub fn last(&self) -> Option<i64> {
102        self.lock().last
103    }
104
105    /// How many values have been handed out.
106    #[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    /// `nextval`: the counter's value, moving the counter on by the increment.
116    ///
117    /// # Errors
118    ///
119    /// A sequence without `CYCLE` that has run past an end. The counter has moved all the same,
120    /// which is what the pin does too.
121    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    /// `currval`: the value last handed out.
161    ///
162    /// # Errors
163    ///
164    /// None has been yet.
165    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    /// `setval`: puts the counter at `value`, and with `called` hands that value out as `nextval`
172    /// would, so the next call gives the one after it.
173    ///
174    /// # Errors
175    ///
176    /// A value outside the sequence's bounds.
177    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
195/// The counter with this number.
196///
197/// # Errors
198///
199/// None is alive with it, which means the sequence was dropped after the statement was bound.
200pub 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}