use std::collections::HashMap;
use std::sync::atomic::{AtomicBool, AtomicI64, Ordering};
use std::sync::{Arc, OnceLock, Weak};
use parking_lot::Mutex;
use crate::storage::StorageEngine;
use crate::storage::{PersistedSeqState, PersistedSequence};
use crate::{Error, Result};
struct SeqRuntime {
increment: i64,
min_value: i64,
max_value: i64,
cache: i64,
cycle: bool,
start: i64,
not_called_target: i64,
next: AtomicI64,
block_end: AtomicI64,
is_called: AtomicBool,
last_served: AtomicI64,
refill: Mutex<()>,
volatile: bool,
}
static STORE: OnceLock<Mutex<HashMap<String, Arc<SeqRuntime>>>> = OnceLock::new();
static PERSIST: OnceLock<Mutex<Option<Weak<StorageEngine>>>> = OnceLock::new();
fn store() -> &'static Mutex<HashMap<String, Arc<SeqRuntime>>> {
STORE.get_or_init(|| Mutex::new(HashMap::new()))
}
pub fn install_persistence(engine: &Arc<StorageEngine>) {
*PERSIST.get_or_init(|| Mutex::new(None)).lock() = Some(Arc::downgrade(engine));
}
pub fn persist_handle() -> Option<Arc<StorageEngine>> {
PERSIST.get().and_then(|m| m.lock().clone()).and_then(|w| w.upgrade())
}
pub fn invalidate_cache(name: &str) {
if let Some(m) = STORE.get() {
m.lock().remove(name);
}
}
pub fn warm_load(engine: &StorageEngine) -> Result<()> {
let defs = engine.catalog().list_sequences()?;
let mut guard = store().lock();
for def in defs {
let st = engine.catalog().get_sequence_state(&def.name)?;
let rt = SeqRuntime::from_persisted(&def, st, false);
guard.insert(def.name.clone(), Arc::new(rt));
}
Ok(())
}
impl SeqRuntime {
fn seed_state(def: &PersistedSequence) -> PersistedSeqState {
PersistedSeqState {
last_reserved: def.start_value,
is_called: false,
}
}
fn from_persisted(def: &PersistedSequence, state: Option<PersistedSeqState>, volatile: bool) -> Self {
let st = state.unwrap_or_else(|| Self::seed_state(def));
let increment = def.increment_by;
let (next, block_end) = if st.is_called {
let parked = if increment > 0 {
st.last_reserved.saturating_add(1)
} else {
st.last_reserved.saturating_sub(1)
};
(parked, st.last_reserved)
} else {
let empty_end = if increment > 0 {
st.last_reserved.saturating_sub(1)
} else {
st.last_reserved.saturating_add(1)
};
(st.last_reserved, empty_end)
};
SeqRuntime {
increment,
min_value: def.min_value,
max_value: def.max_value,
cache: def.cache.max(1),
cycle: def.cycle,
start: def.start_value,
not_called_target: st.last_reserved,
next: AtomicI64::new(next),
block_end: AtomicI64::new(block_end),
is_called: AtomicBool::new(st.is_called),
last_served: AtomicI64::new(i64::MIN),
refill: Mutex::new(()),
volatile,
}
}
#[inline]
fn in_block(&self, cur: i64, end: i64) -> bool {
if self.increment > 0 {
cur <= end
} else {
cur >= end
}
}
}
fn checked_step(name: &str, cur: i64, incr: i64, min: i64, max: i64, cycle: bool) -> Result<i64> {
match cur.checked_add(incr) {
Some(n) if incr > 0 && n <= max => Ok(n),
Some(n) if incr < 0 && n >= min => Ok(n),
Some(n) if incr == 0 => Ok(n),
_ => {
if cycle {
Ok(if incr > 0 { min } else { max })
} else if incr > 0 {
Err(Error::query_execution(format!(
"nextval: reached maximum value of sequence \"{}\" ({})",
name, max
)))
} else {
Err(Error::query_execution(format!(
"nextval: reached minimum value of sequence \"{}\" ({})",
name, min
)))
}
}
}
}
fn clamp_block_end(first: i64, incr: i64, cache: i64, min: i64, max: i64) -> i64 {
let span = (incr as i128) * ((cache.max(1) - 1) as i128);
let raw = first as i128 + span;
let clamped = raw.clamp(min as i128, max as i128);
clamped as i64
}
fn runtime_for(name: &str) -> Result<Arc<SeqRuntime>> {
{
let guard = store().lock();
if let Some(rt) = guard.get(name) {
return Ok(Arc::clone(rt));
}
}
let rt = match persist_handle() {
Some(engine) => match engine.catalog().get_sequence(name)? {
Some(def) => {
let st = engine.catalog().get_sequence_state(name)?;
Arc::new(SeqRuntime::from_persisted(&def, st, false))
}
None => {
let def = PersistedSequence::default_named(name);
let mut guard = store().lock();
if let Some(existing) = guard.get(name) {
return Ok(Arc::clone(existing));
}
engine.catalog().save_sequence(&def)?;
let st = match engine.catalog().get_sequence_state(name)? {
Some(existing) => existing,
None => {
let seed = SeqRuntime::seed_state(&def);
engine.catalog().save_sequence_state(name, &seed)?;
seed
}
};
let rt = Arc::new(SeqRuntime::from_persisted(&def, Some(st), false));
guard.insert(name.to_string(), Arc::clone(&rt));
return Ok(rt);
}
},
None => {
let def = PersistedSequence::default_named(name);
Arc::new(SeqRuntime::from_persisted(&def, None, true))
}
};
let mut guard = store().lock();
let entry = guard.entry(name.to_string()).or_insert_with(|| Arc::clone(&rt));
Ok(Arc::clone(entry))
}
fn persist_high_water(rt: &SeqRuntime, name: &str, last_reserved: i64) -> Result<()> {
if rt.volatile {
return Ok(());
}
let engine = persist_handle().ok_or_else(|| Error::query_execution("nextval requires storage context"))?;
engine.flush_sequence_state(
name,
PersistedSeqState {
last_reserved,
is_called: true,
},
)
}
pub fn try_nextval(name: &str) -> Result<i64> {
let rt = runtime_for(name)?;
loop {
let cur = rt.next.load(Ordering::Acquire);
let end = rt.block_end.load(Ordering::Acquire);
if !rt.in_block(cur, end) {
break; }
let stepped = match cur.checked_add(rt.increment) {
Some(n) => n,
None => {
if rt.increment > 0 {
end.saturating_add(1)
} else {
end.saturating_sub(1)
}
}
};
if rt
.next
.compare_exchange(cur, stepped, Ordering::AcqRel, Ordering::Acquire)
.is_ok()
{
rt.is_called.store(true, Ordering::Relaxed);
rt.last_served.store(cur, Ordering::Relaxed);
return Ok(cur);
}
}
let _g = rt.refill.lock();
{
let cur = rt.next.load(Ordering::Acquire);
let end = rt.block_end.load(Ordering::Acquire);
if rt.in_block(cur, end) {
drop(_g);
return try_nextval(name);
}
}
let in_mem_high = rt.block_end.load(Ordering::Acquire);
let in_mem_called = rt.is_called.load(Ordering::Acquire);
let durable: Option<PersistedSeqState> = if rt.volatile {
None
} else {
match persist_handle() {
Some(engine) => engine.catalog().get_sequence_state(name)?,
None => None,
}
};
let durable_called = durable.map(|d| d.is_called).unwrap_or(false);
let base_called = durable_called || in_mem_called;
let first = if !base_called {
match durable {
Some(d) if !d.is_called => d.last_reserved,
_ => rt.not_called_target,
}
} else {
let durable_high = durable.map(|d| d.last_reserved).unwrap_or(in_mem_high);
let high = std::cmp::max(durable_high, in_mem_high);
checked_step(name, high, rt.increment, rt.min_value, rt.max_value, rt.cycle)?
};
let last = clamp_block_end(first, rt.increment, rt.cache, rt.min_value, rt.max_value);
persist_high_water(&rt, name, last)?;
let next_after = match first.checked_add(rt.increment) {
Some(n) => n,
None => {
if rt.increment > 0 {
last.saturating_add(1)
} else {
last.saturating_sub(1)
}
}
};
rt.next.store(next_after, Ordering::Release);
rt.is_called.store(true, Ordering::Relaxed);
rt.last_served.store(first, Ordering::Relaxed);
rt.block_end.store(last, Ordering::Release);
drop(_g);
Ok(first)
}
pub fn try_currval(name: &str) -> Result<i64> {
let rt = match STORE.get().and_then(|m| m.lock().get(name).cloned()) {
Some(rt) => rt,
None => return Ok(0),
};
if !rt.is_called.load(Ordering::Acquire) {
return Ok(0);
}
let v = rt.last_served.load(Ordering::Relaxed);
if v != i64::MIN {
return Ok(v);
}
let next = rt.next.load(Ordering::Acquire);
let last = next.checked_sub(rt.increment).unwrap_or(next);
Ok(last)
}
pub fn peek_last_served(name: &str) -> Option<i64> {
let guard = STORE.get()?.lock();
let rt = guard.get(name)?;
let v = rt.last_served.load(Ordering::Relaxed);
if v == i64::MIN {
None
} else {
Some(v)
}
}
pub fn try_setval(name: &str, value: i64, is_called: bool) -> Result<i64> {
let rt = runtime_for(name)?;
if value < rt.min_value || value > rt.max_value {
return Err(Error::query_execution(format!(
"setval: value {} is out of bounds for sequence \"{}\" (min {}, max {})",
value, name, rt.min_value, rt.max_value
)));
}
let _g = rt.refill.lock();
if !rt.volatile {
let engine = persist_handle().ok_or_else(|| Error::query_execution("setval requires storage context"))?;
engine.flush_sequence_state(
name,
PersistedSeqState {
last_reserved: value,
is_called,
},
)?;
drop(_g);
invalidate_cache(name);
return Ok(value);
}
if is_called {
let nxt = checked_step(name, value, rt.increment, rt.min_value, rt.max_value, true).unwrap_or(value);
rt.block_end.store(value, Ordering::Release);
rt.next.store(nxt, Ordering::Release);
} else {
rt.block_end.store(value, Ordering::Release);
rt.next.store(value, Ordering::Release);
}
rt.is_called.store(is_called, Ordering::Release);
drop(_g);
Ok(value)
}
pub fn install_runtime(def: &PersistedSequence, state: Option<PersistedSeqState>) {
let rt = Arc::new(SeqRuntime::from_persisted(def, state, persist_handle().is_none()));
store().lock().insert(def.name.clone(), rt);
}
pub fn create_sequence(name: &str, if_not_exists: bool, start: Option<i64>, increment: Option<i64>) {
let increment = match increment.unwrap_or(1) {
0 => 1,
n => n,
};
let start = start.unwrap_or(1);
if if_not_exists {
if store().lock().contains_key(name) {
return;
}
if let Some(engine) = persist_handle() {
if engine.catalog().sequence_exists(name).unwrap_or(false) {
return;
}
}
}
let (mut min_value, max_value) = if increment > 0 {
(1, PersistedSequence::BIGINT_MAX)
} else {
(PersistedSequence::BIGINT_MIN, -1)
};
if increment > 0 && start < min_value {
min_value = start;
}
let def = PersistedSequence {
name: name.to_string(),
data_type: "bigint".to_string(),
start_value: start,
increment_by: increment,
min_value,
max_value,
cache: 1,
cycle: false,
owned_by_table: None,
owned_by_column: None,
};
if let Some(engine) = persist_handle() {
let _ = engine.catalog().save_sequence(&def);
let _ = engine
.catalog()
.save_sequence_state(name, &SeqRuntime::seed_state(&def));
}
install_runtime(&def, None);
}
pub fn nextval(name: &str) -> i64 {
match try_nextval(name) {
Ok(v) => v,
Err(_) => {
match runtime_for(name) {
Ok(rt) => {
if rt.increment >= 0 {
rt.max_value
} else {
rt.min_value
}
}
Err(_) => 0,
}
}
}
}
pub fn currval(name: &str) -> i64 {
try_currval(name).unwrap_or(0)
}
pub fn setval(name: &str, value: i64, is_called: bool) -> i64 {
let _ = try_setval(name, value, is_called);
value
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_sequence_starts_at_one() {
create_sequence("seq_default", false, None, None);
assert_eq!(nextval("seq_default"), 1);
assert_eq!(nextval("seq_default"), 2);
assert_eq!(currval("seq_default"), 2);
}
#[test]
fn honors_start_and_increment() {
create_sequence("seq_si", false, Some(100), Some(10));
assert_eq!(nextval("seq_si"), 100);
assert_eq!(nextval("seq_si"), 110);
assert_eq!(nextval("seq_si"), 120);
}
#[test]
fn setval_preserves_increment() {
create_sequence("seq_sv", false, Some(1), Some(5));
assert_eq!(nextval("seq_sv"), 1);
setval("seq_sv", 50, true);
assert_eq!(nextval("seq_sv"), 55);
}
#[test]
fn setval_is_called_false_makes_next_nextval_equal_value() {
create_sequence("seq_sv_false", false, Some(1), Some(5));
assert_eq!(nextval("seq_sv_false"), 1);
setval("seq_sv_false", 300, false);
assert_eq!(nextval("seq_sv_false"), 300);
assert_eq!(nextval("seq_sv_false"), 305);
}
#[test]
fn unknown_sequence_auto_vivifies_at_one() {
assert_eq!(nextval("seq_never_created_xyz"), 1);
}
#[test]
fn cache_serves_contiguous_values_in_one_window() {
let def = PersistedSequence {
name: "seq_cache".into(),
data_type: "bigint".into(),
start_value: 1,
increment_by: 1,
min_value: 1,
max_value: PersistedSequence::BIGINT_MAX,
cache: 8,
cycle: false,
owned_by_table: None,
owned_by_column: None,
};
install_runtime(&def, None);
for expected in 1..=20 {
assert_eq!(try_nextval("seq_cache").unwrap(), expected);
}
}
#[test]
fn ascending_maxvalue_no_cycle_errors() {
let def = PersistedSequence {
name: "seq_maxnc".into(),
data_type: "bigint".into(),
start_value: 1,
increment_by: 1,
min_value: 1,
max_value: 3,
cache: 1,
cycle: false,
owned_by_table: None,
owned_by_column: None,
};
install_runtime(&def, None);
assert_eq!(try_nextval("seq_maxnc").unwrap(), 1);
assert_eq!(try_nextval("seq_maxnc").unwrap(), 2);
assert_eq!(try_nextval("seq_maxnc").unwrap(), 3);
let err = try_nextval("seq_maxnc").unwrap_err();
assert!(err.to_string().contains("reached maximum value"), "{}", err);
}
#[test]
fn ascending_maxvalue_cycle_wraps_to_min() {
let def = PersistedSequence {
name: "seq_maxcy".into(),
data_type: "bigint".into(),
start_value: 1,
increment_by: 1,
min_value: 1,
max_value: 3,
cache: 1,
cycle: true,
owned_by_table: None,
owned_by_column: None,
};
install_runtime(&def, None);
assert_eq!(try_nextval("seq_maxcy").unwrap(), 1);
assert_eq!(try_nextval("seq_maxcy").unwrap(), 2);
assert_eq!(try_nextval("seq_maxcy").unwrap(), 3);
assert_eq!(try_nextval("seq_maxcy").unwrap(), 1);
}
#[test]
fn descending_minvalue_no_cycle_errors() {
let def = PersistedSequence {
name: "seq_desc".into(),
data_type: "bigint".into(),
start_value: -1,
increment_by: -1,
min_value: -3,
max_value: -1,
cache: 1,
cycle: false,
owned_by_table: None,
owned_by_column: None,
};
install_runtime(&def, None);
assert_eq!(try_nextval("seq_desc").unwrap(), -1);
assert_eq!(try_nextval("seq_desc").unwrap(), -2);
assert_eq!(try_nextval("seq_desc").unwrap(), -3);
let err = try_nextval("seq_desc").unwrap_err();
assert!(err.to_string().contains("reached minimum value"), "{}", err);
}
#[test]
fn overflow_near_i64_max_errors_not_panics() {
let def = PersistedSequence {
name: "seq_ovf".into(),
data_type: "bigint".into(),
start_value: i64::MAX - 1,
increment_by: 10,
min_value: 1,
max_value: PersistedSequence::BIGINT_MAX,
cache: 1,
cycle: false,
owned_by_table: None,
owned_by_column: None,
};
install_runtime(&def, None);
assert_eq!(try_nextval("seq_ovf").unwrap(), i64::MAX - 1);
let err = try_nextval("seq_ovf").unwrap_err();
assert!(err.to_string().contains("reached maximum value"), "{}", err);
}
#[test]
fn setval_out_of_bounds_errors() {
let def = PersistedSequence {
name: "seq_svb".into(),
data_type: "bigint".into(),
start_value: 1,
increment_by: 1,
min_value: 1,
max_value: 100,
cache: 1,
cycle: false,
owned_by_table: None,
owned_by_column: None,
};
install_runtime(&def, None);
assert!(try_setval("seq_svb", 999, true).is_err());
assert_eq!(try_setval("seq_svb", 50, true).unwrap(), 50);
}
#[test]
fn cache_never_exceeds_max_bound() {
let def = PersistedSequence {
name: "seq_clamp".into(),
data_type: "bigint".into(),
start_value: 1,
increment_by: 1,
min_value: 1,
max_value: 5,
cache: 100,
cycle: false,
owned_by_table: None,
owned_by_column: None,
};
install_runtime(&def, None);
for expected in 1..=5 {
assert_eq!(try_nextval("seq_clamp").unwrap(), expected);
}
assert!(try_nextval("seq_clamp")
.unwrap_err()
.to_string()
.contains("reached maximum value"));
}
#[test]
fn clamp_block_end_stays_in_range() {
assert_eq!(clamp_block_end(1, 1, 100, 1, 5), 5);
assert_eq!(clamp_block_end(1, 1, 8, 1, i64::MAX), 8);
assert_eq!(clamp_block_end(-1, -1, 100, -5, -1), -5);
assert_eq!(clamp_block_end(i64::MAX - 1, 10, 1000, 1, i64::MAX), i64::MAX);
}
}