use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, LazyLock, Mutex, Weak};
use crate::{Error, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub struct Options {
pub increment: i64,
pub min: i64,
pub max: i64,
pub start: i64,
pub cycle: bool,
}
#[derive(Debug, Clone, Copy)]
struct State {
counter: i64,
last: Option<i64>,
uses: u64,
}
#[derive(Debug)]
pub struct Counter {
id: u64,
name: String,
options: Options,
state: Mutex<State>,
}
static COUNTERS: LazyLock<Mutex<HashMap<u64, Weak<Counter>>>> =
LazyLock::new(|| Mutex::new(HashMap::new()));
static NEXT: AtomicU64 = AtomicU64::new(1);
impl Counter {
#[must_use]
pub fn register(name: &str, options: Options) -> Arc<Self> {
let id = NEXT.fetch_add(1, Ordering::Relaxed);
let counter = Arc::new(Self {
id,
name: name.to_owned(),
options,
state: Mutex::new(State { counter: options.start, last: None, uses: 0 }),
});
let mut all = COUNTERS.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
all.retain(|_, counter| counter.strong_count() > 0);
all.insert(id, Arc::downgrade(&counter));
counter
}
#[must_use]
pub fn id(&self) -> u64 {
self.id
}
#[must_use]
pub fn options(&self) -> Options {
self.options
}
#[must_use]
pub fn counter(&self) -> i64 {
self.lock().counter
}
#[must_use]
pub fn last(&self) -> Option<i64> {
self.lock().last
}
#[must_use]
pub fn uses(&self) -> u64 {
self.lock().uses
}
fn lock(&self) -> std::sync::MutexGuard<'_, State> {
self.state.lock().unwrap_or_else(std::sync::PoisonError::into_inner)
}
pub fn next(&self) -> Result<i64> {
let mut state = self.lock();
self.advance(&mut state)
}
fn advance(&self, state: &mut State) -> Result<i64> {
let Options { increment, min, max, cycle, .. } = self.options;
let result = state.counter;
let (moved, overflow) = state.counter.overflowing_add(increment);
if !overflow {
state.counter = moved;
}
if cycle {
if overflow {
state.counter = if increment < 0 { max } else { min };
} else if state.counter < min {
state.counter = max;
} else if state.counter > max {
state.counter = min;
}
} else {
if result < min || (overflow && increment < 0) {
return Err(Error::sequence(format!(
"nextval: reached minimum value of sequence \"{}\" ({min})",
self.name
)));
}
if result > max || overflow {
return Err(Error::sequence(format!(
"nextval: reached maximum value of sequence \"\"{}\"\" ({max})",
self.name
)));
}
}
state.last = Some(result);
state.uses += 1;
Ok(result)
}
pub fn current(&self) -> Result<i64> {
self.lock()
.last
.ok_or_else(|| Error::sequence("currval: sequence is not yet defined in this session"))
}
pub fn set(&self, value: i64, called: bool) -> Result<i64> {
let Options { min, max, .. } = self.options;
if value < min || value > max {
return Err(Error::sequence(format!(
"setval: value {value} is out of bounds for sequence \"{}\" ({min}..{max})",
self.name
)));
}
let mut state = self.lock();
state.counter = value;
if !called {
state.uses += 1;
return Ok(value);
}
self.advance(&mut state)
}
}
pub fn lookup(id: u64) -> Result<Arc<Counter>> {
let all = COUNTERS.lock().unwrap_or_else(std::sync::PoisonError::into_inner);
all.get(&id)
.and_then(Weak::upgrade)
.ok_or_else(|| Error::catalog("The sequence this statement uses no longer exists"))
}
#[cfg(test)]
mod tests {
use super::*;
fn options(increment: i64, min: i64, max: i64, start: i64, cycle: bool) -> Options {
Options { increment, min, max, start, cycle }
}
#[test]
fn a_counter_counts_and_stops_at_its_maximum() {
let seq = Counter::register("seq", options(1, 1, 2, 1, false));
assert_eq!(seq.next().unwrap(), 1);
assert_eq!(seq.next().unwrap(), 2);
let error = seq.next().unwrap_err().to_string();
assert!(error.contains("reached maximum value of sequence \"\"seq\"\" (2)"), "{error}");
assert_eq!(seq.current().unwrap(), 2);
assert!(Arc::ptr_eq(&lookup(seq.id()).unwrap(), &seq));
}
#[test]
fn a_cycle_starts_again_from_the_other_end() {
let seq = Counter::register("seq", options(-1, 1, 3, 2, true));
let got: Vec<i64> = (0..4).map(|_| seq.next().unwrap()).collect();
assert_eq!(got, [2, 1, 3, 2]);
}
#[test]
fn setval_hands_out_the_value_when_called() {
let seq = Counter::register("seq", options(1, 1, i64::MAX, 1, false));
assert!(seq.current().is_err());
assert_eq!(seq.set(10, true).unwrap(), 10);
assert_eq!(seq.next().unwrap(), 11);
assert_eq!(seq.set(5, false).unwrap(), 5);
assert_eq!(seq.current().unwrap(), 11);
assert_eq!(seq.next().unwrap(), 5);
assert!(seq.set(0, true).is_err());
}
}