use super::engine::KvEngine;
use super::engine_atomic_compute as compute;
use super::engine_helpers::{expiry_key, table_key};
use super::entry::NO_EXPIRY;
use super::hash_table::KvHashTable;
pub struct CasResult {
pub success: bool,
pub current_value: Option<Vec<u8>>,
}
#[derive(Debug)]
pub enum AtomicError {
TypeMismatch { detail: String },
Overflow,
Encode { detail: String },
}
#[derive(Clone, Copy)]
pub struct AtomicKeyCtx<'a> {
pub database_id: u64,
pub tenant_id: u64,
pub collection: &'a str,
pub key: &'a [u8],
pub now_ms: u64,
pub surrogate: nodedb_types::Surrogate,
}
impl KvEngine {
pub fn incr(
&mut self,
ctx: AtomicKeyCtx<'_>,
delta: i64,
ttl_ms: u64,
) -> Result<i64, AtomicError> {
self.incr_resolved(ctx, delta, ttl_ms, None)
}
pub fn incr_with_absolute_expiry(
&mut self,
ctx: AtomicKeyCtx<'_>,
delta: i64,
ttl_ms: u64,
expire_at_ms: u64,
) -> Result<i64, AtomicError> {
self.incr_resolved(ctx, delta, ttl_ms, Some(expire_at_ms))
}
fn incr_resolved(
&mut self,
ctx: AtomicKeyCtx<'_>,
delta: i64,
ttl_ms: u64,
expire_override: Option<u64>,
) -> Result<i64, AtomicError> {
let tkey = table_key(ctx.database_id, ctx.tenant_id, ctx.collection);
let table = self.ensure_table(tkey, ctx.tenant_id, ctx.collection);
let current = table.get(ctx.key, ctx.now_ms).map(|v| v.to_vec());
let (new_i64, new_bytes) = compute::incr(current.as_deref(), delta)?;
self.atomic_put(
ctx,
tkey,
&new_bytes,
ttl_ms,
current.is_none(),
expire_override,
);
Ok(new_i64)
}
pub fn incr_float(&mut self, ctx: AtomicKeyCtx<'_>, delta: f64) -> Result<f64, AtomicError> {
let tkey = table_key(ctx.database_id, ctx.tenant_id, ctx.collection);
let table = self.ensure_table(tkey, ctx.tenant_id, ctx.collection);
let current = table.get(ctx.key, ctx.now_ms).map(|v| v.to_vec());
let (new_f64, new_bytes) = compute::incr_float(current.as_deref(), delta)?;
self.atomic_put(ctx, tkey, &new_bytes, 0, current.is_none(), None);
Ok(new_f64)
}
pub fn cas(&mut self, ctx: AtomicKeyCtx<'_>, expected: &[u8], new_value: &[u8]) -> CasResult {
let tkey = table_key(ctx.database_id, ctx.tenant_id, ctx.collection);
let table = self.ensure_table(tkey, ctx.tenant_id, ctx.collection);
let current = table.get(ctx.key, ctx.now_ms).map(|v| v.to_vec());
let (matches, write_bytes) = compute::cas(current.as_deref(), expected, new_value);
if matches {
self.atomic_put(ctx, tkey, &write_bytes, 0, current.is_none(), None);
CasResult {
success: true,
current_value: current,
}
} else {
CasResult {
success: false,
current_value: current,
}
}
}
pub fn getset(&mut self, ctx: AtomicKeyCtx<'_>, new_value: &[u8]) -> Option<Vec<u8>> {
let tkey = table_key(ctx.database_id, ctx.tenant_id, ctx.collection);
let table = self.ensure_table(tkey, ctx.tenant_id, ctx.collection);
let old = table.get(ctx.key, ctx.now_ms).map(|v| v.to_vec());
let write_bytes = compute::getset(old.as_deref(), new_value);
self.atomic_put(ctx, tkey, &write_bytes, 0, old.is_none(), None);
old
}
fn ensure_table(&mut self, tkey: u64, tenant_id: u64, collection: &str) -> &mut KvHashTable {
self.hash_to_tenant.entry(tkey).or_insert(tenant_id);
self.hash_to_collection
.entry(tkey)
.or_insert_with(|| collection.to_string());
let default_capacity = self.default_capacity;
let load_factor_threshold = self.load_factor_threshold;
let rehash_batch_size = self.rehash_batch_size;
let inline_threshold = self.inline_threshold;
self.tables.entry(tkey).or_insert_with(|| {
KvHashTable::new(
default_capacity,
load_factor_threshold,
rehash_batch_size,
inline_threshold,
)
})
}
fn atomic_put(
&mut self,
ctx: AtomicKeyCtx<'_>,
tkey: u64,
value: &[u8],
ttl_ms: u64,
is_new_key: bool,
expire_override: Option<u64>,
) {
let AtomicKeyCtx {
database_id,
tenant_id,
collection,
key,
now_ms,
surrogate,
} = ctx;
let old_meta = if is_new_key {
None
} else {
self.tables.get(&tkey).and_then(|t| t.get_entry_meta(key))
};
let expire_at = if ttl_ms > 0 {
expire_override.unwrap_or(now_ms + ttl_ms)
} else if let Some(ref meta) = old_meta {
meta.expire_at_ms
} else {
NO_EXPIRY
};
if let Some(ref meta) = old_meta
&& meta.has_ttl
{
let composite = expiry_key(database_id, tenant_id, collection, key);
self.expiry.cancel(&composite, meta.expire_at_ms);
}
let old_fields =
if !is_new_key && self.indexes.get(&tkey).is_some_and(|idx| !idx.is_empty()) {
self.tables
.get(&tkey)
.and_then(|t| t.get(key, now_ms))
.map(|old_val| {
super::engine_helpers::extract_all_field_values_from_msgpack(old_val)
})
} else {
None
};
let default_capacity = self.default_capacity;
let load_factor_threshold = self.load_factor_threshold;
let rehash_batch_size = self.rehash_batch_size;
let inline_threshold = self.inline_threshold;
let table = self.tables.entry(tkey).or_insert_with(|| {
KvHashTable::new(
default_capacity,
load_factor_threshold,
rehash_batch_size,
inline_threshold,
)
});
table.put(key, value, expire_at, surrogate);
if expire_at != NO_EXPIRY {
let composite = expiry_key(database_id, tenant_id, collection, key);
self.expiry.insert(composite, expire_at);
}
if self.indexes.get(&tkey).is_some_and(|idx| !idx.is_empty()) {
let old_refs: Option<Vec<(&str, &[u8])>> = old_fields.as_ref().map(|fields| {
fields
.iter()
.map(|(k, v)| (k.as_str(), v.as_slice()))
.collect()
});
let new_fields = super::engine_helpers::extract_all_field_values_from_msgpack(value);
let new_refs: Vec<(&str, &[u8])> = new_fields
.iter()
.map(|(k, v)| (k.as_str(), v.as_slice()))
.collect();
if let Some(idx_set) = self.indexes.get_mut(&tkey) {
idx_set.on_put(key, &new_refs, old_refs.as_deref());
}
}
}
}
#[cfg(test)]
mod tests {
use nodedb_types::Surrogate;
use super::super::engine_write::KvPutParams;
use super::*;
fn make_engine() -> KvEngine {
KvEngine::new(1000, 16, 0.75, 4, 64, 1000, 1024)
}
fn ctx<'a>(collection: &'a str, key: &'a [u8]) -> AtomicKeyCtx<'a> {
AtomicKeyCtx {
database_id: 0,
tenant_id: 1,
collection,
key,
now_ms: 1000,
surrogate: Surrogate::ZERO,
}
}
#[test]
fn incr_new_key() {
let mut engine = make_engine();
let result = engine.incr(ctx("counters", b"hits"), 10, 0);
assert_eq!(result.unwrap(), 10);
}
#[test]
fn incr_existing_key() {
let mut engine = make_engine();
engine.incr(ctx("counters", b"hits"), 10, 0).unwrap();
let result = engine.incr(ctx("counters", b"hits"), 5, 0);
assert_eq!(result.unwrap(), 15);
}
#[test]
fn incr_negative_delta() {
let mut engine = make_engine();
engine.incr(ctx("counters", b"gold"), 100, 0).unwrap();
let result = engine.incr(ctx("counters", b"gold"), -30, 0);
assert_eq!(result.unwrap(), 70);
}
#[test]
fn incr_overflow() {
let mut engine = make_engine();
let bytes = zerompk::to_msgpack_vec(&i64::MAX).unwrap();
engine.put(KvPutParams {
database_id: 0,
tenant_id: 1,
collection: "counters",
key: b"max",
value: &bytes,
ttl_ms: 0,
now_ms: 1000,
surrogate: Surrogate::ZERO,
});
let result = engine.incr(ctx("counters", b"max"), 1, 0);
assert!(matches!(result, Err(AtomicError::Overflow)));
}
#[test]
fn incr_type_mismatch() {
let mut engine = make_engine();
let bytes = zerompk::to_msgpack_vec(&"hello").unwrap();
engine.put(KvPutParams {
database_id: 0,
tenant_id: 1,
collection: "counters",
key: b"str",
value: &bytes,
ttl_ms: 0,
now_ms: 1000,
surrogate: Surrogate::ZERO,
});
let result = engine.incr(ctx("counters", b"str"), 1, 0);
assert!(matches!(result, Err(AtomicError::TypeMismatch { .. })));
}
#[test]
fn incr_with_ttl_new_key() {
let mut engine = make_engine();
engine
.incr(ctx("counters", b"daily"), 1, 86_400_000)
.unwrap();
let ttl = engine.get_ttl_ms(0, 1, "counters", b"daily", 1000);
assert!(ttl.is_some());
assert!(ttl.unwrap() > 0);
}
#[test]
fn incr_preserves_ttl_when_zero() {
let mut engine = make_engine();
let bytes = zerompk::to_msgpack_vec(&50i64).unwrap();
engine.put(KvPutParams {
database_id: 0,
tenant_id: 1,
collection: "counters",
key: b"temp",
value: &bytes,
ttl_ms: 5000,
now_ms: 1000,
surrogate: Surrogate::ZERO,
});
engine.incr(ctx("counters", b"temp"), 10, 0).unwrap();
let ttl = engine.get_ttl_ms(0, 1, "counters", b"temp", 1000);
assert!(ttl.is_some());
assert!(ttl.unwrap() > 0);
}
#[test]
fn incr_with_absolute_expiry_installs_recorded_instant_not_now_plus_ttl() {
let mut engine = make_engine();
engine
.incr_with_absolute_expiry(ctx("counters", b"daily"), 1, 5_000, 1_000_000)
.unwrap();
let ttl = engine.get_ttl_ms(0, 1, "counters", b"daily", 1000);
assert_eq!(
ttl,
Some(1_000_000 - 1000),
"must install the caller-supplied absolute instant verbatim, not now_ms + ttl_ms"
);
}
#[test]
fn incr_with_absolute_expiry_and_zero_ttl_still_preserves_existing_expiry() {
let mut engine = make_engine();
let bytes = zerompk::to_msgpack_vec(&50i64).unwrap();
engine.put(KvPutParams {
database_id: 0,
tenant_id: 1,
collection: "counters",
key: b"temp",
value: &bytes,
ttl_ms: 5000,
now_ms: 1000,
surrogate: Surrogate::ZERO,
});
let ttl_before = engine.get_ttl_ms(0, 1, "counters", b"temp", 1000);
engine
.incr_with_absolute_expiry(ctx("counters", b"temp"), 10, 0, 999_999_999)
.unwrap();
let ttl_after = engine.get_ttl_ms(0, 1, "counters", b"temp", 1000);
assert_eq!(
ttl_before, ttl_after,
"ttl_ms == 0 must preserve the existing expiry, ignoring expire_override"
);
}
#[test]
fn incr_float_new_key() {
let mut engine = make_engine();
let result = engine.incr_float(ctx("scores", b"dmg"), 3.125);
assert!((result.unwrap() - 3.125).abs() < f64::EPSILON);
}
#[test]
fn incr_float_existing() {
let mut engine = make_engine();
engine.incr_float(ctx("scores", b"dmg"), 3.0).unwrap();
let result = engine.incr_float(ctx("scores", b"dmg"), 1.5);
assert!((result.unwrap() - 4.5).abs() < f64::EPSILON);
}
#[test]
fn incr_float_infinity_rejected() {
let mut engine = make_engine();
let bytes = zerompk::to_msgpack_vec(&f64::MAX).unwrap();
engine.put(KvPutParams {
database_id: 0,
tenant_id: 1,
collection: "scores",
key: b"big",
value: &bytes,
ttl_ms: 0,
now_ms: 1000,
surrogate: Surrogate::ZERO,
});
let result = engine.incr_float(ctx("scores", b"big"), f64::MAX);
assert!(matches!(result, Err(AtomicError::Overflow)));
}
#[test]
fn cas_create_if_not_exists() {
let mut engine = make_engine();
let result = engine.cas(ctx("state", b"player1"), b"", b"idle");
assert!(result.success);
assert!(result.current_value.is_none());
let val = engine.get(0, 1, "state", b"player1", 1000);
assert_eq!(val.as_deref(), Some(b"idle".as_slice()));
}
#[test]
fn cas_success() {
let mut engine = make_engine();
engine.put(KvPutParams {
database_id: 0,
tenant_id: 1,
collection: "state",
key: b"p1",
value: b"idle",
ttl_ms: 0,
now_ms: 1000,
surrogate: Surrogate::ZERO,
});
let result = engine.cas(ctx("state", b"p1"), b"idle", b"in_match");
assert!(result.success);
assert_eq!(result.current_value.as_deref(), Some(b"idle".as_slice()));
let val = engine.get(0, 1, "state", b"p1", 1000);
assert_eq!(val.as_deref(), Some(b"in_match".as_slice()));
}
#[test]
fn cas_failure() {
let mut engine = make_engine();
engine.put(KvPutParams {
database_id: 0,
tenant_id: 1,
collection: "state",
key: b"p1",
value: b"fighting",
ttl_ms: 0,
now_ms: 1000,
surrogate: Surrogate::ZERO,
});
let result = engine.cas(ctx("state", b"p1"), b"idle", b"in_match");
assert!(!result.success);
assert_eq!(
result.current_value.as_deref(),
Some(b"fighting".as_slice())
);
let val = engine.get(0, 1, "state", b"p1", 1000);
assert_eq!(val.as_deref(), Some(b"fighting".as_slice()));
}
#[test]
fn getset_new_key() {
let mut engine = make_engine();
let old = engine.getset(ctx("session", b"tok"), b"new-token");
assert!(old.is_none());
let val = engine.get(0, 1, "session", b"tok", 1000);
assert_eq!(val.as_deref(), Some(b"new-token".as_slice()));
}
#[test]
fn getset_existing_key() {
let mut engine = make_engine();
engine.put(KvPutParams {
database_id: 0,
tenant_id: 1,
collection: "session",
key: b"tok",
value: b"old-token",
ttl_ms: 0,
now_ms: 1000,
surrogate: Surrogate::ZERO,
});
let old = engine.getset(ctx("session", b"tok"), b"new-token");
assert_eq!(old.as_deref(), Some(b"old-token".as_slice()));
let val = engine.get(0, 1, "session", b"tok", 1000);
assert_eq!(val.as_deref(), Some(b"new-token".as_slice()));
}
}