use crate::{KevyError, KevyResult};
use std::sync::RwLockWriteGuard;
use crate::shard::shard_idx;
use crate::store::{Inner, Store, commit_write, store_err};
use crate::store::ensure_writable;
type ShardUndoEntry = (usize, Vec<u8>, Option<(kevy_store::Value, Option<u64>)>);
pub struct AtomicAllShards<'a> {
pub(crate) guards: Vec<RwLockWriteGuard<'a, Inner>>,
log: Vec<(usize, Vec<Vec<u8>>)>,
undo: Vec<ShardUndoEntry>,
touched: std::collections::HashSet<Vec<u8>>,
#[cfg(feature = "index")]
pub(crate) indexes: std::sync::Arc<crate::ops_index::IndexReg>,
}
impl<'a> AtomicAllShards<'a> {
pub(crate) fn idx(&self, key: &[u8]) -> usize {
shard_idx(key, self.guards.len())
}
fn snap(&mut self, key: &[u8]) {
if self.touched.contains(key) {
return;
}
let i = self.idx(key);
let prior = self.guards[i].store.clone_with_ttl(key);
self.touched.insert(key.to_vec());
self.undo.push((i, key.to_vec(), prior));
}
fn log_arg(&mut self, idx: usize, parts: &[&[u8]]) {
self.log
.push((idx, parts.iter().map(|p| p.to_vec()).collect()));
}
pub fn set(&mut self, key: &[u8], value: &[u8]) -> bool {
self.snap(key);
let i = self.idx(key);
let ok = self.guards[i]
.store
.set(key, value.to_vec(), None, false, false);
self.log_arg(i, &[b"SET", key, value]);
ok
}
pub fn get(&mut self, key: &[u8]) -> KevyResult<Option<Vec<u8>>> {
let i = self.idx(key);
self.guards[i]
.store
.get(key)
.map(|opt| opt.as_deref().map(<[u8]>::to_vec))
.map_err(store_err)
}
pub fn incr(&mut self, key: &[u8]) -> KevyResult<i64> {
self.snap(key);
let i = self.idx(key);
let n = self.guards[i].store.incr_by(key, 1).map_err(store_err)?;
self.log_arg(i, &[b"INCR", key]);
Ok(n)
}
pub fn incr_by(&mut self, key: &[u8], delta: i64) -> KevyResult<i64> {
self.snap(key);
let i = self.idx(key);
let n = self.guards[i].store.incr_by(key, delta).map_err(store_err)?;
let s = format!("{delta}");
self.log_arg(i, &[b"INCRBY", key, s.as_bytes()]);
Ok(n)
}
pub fn hset(&mut self, key: &[u8], pairs: &[(&[u8], &[u8])]) -> KevyResult<usize> {
self.snap(key);
let i = self.idx(key);
let n = self.guards[i]
.store
.hset(key, pairs)
.map_err(store_err)?;
let mut parts: Vec<&[u8]> = Vec::with_capacity(2 + pairs.len() * 2);
parts.push(b"HSET");
parts.push(key);
for (f, v) in pairs {
parts.push(f);
parts.push(v);
}
self.log_arg(i, &parts);
Ok(n)
}
pub fn hget(&mut self, key: &[u8], field: &[u8]) -> KevyResult<Option<Vec<u8>>> {
let i = self.idx(key);
Ok(self.guards[i]
.store
.hget(key, field)
.map_err(store_err)?
.map(<[u8]>::to_vec))
}
pub fn hincrby(&mut self, key: &[u8], field: &[u8], delta: i64) -> KevyResult<i64> {
self.snap(key);
let i = self.idx(key);
let n = self.guards[i]
.store
.hincrby(key, field, delta)
.map_err(store_err)?;
let s = format!("{delta}");
self.log_arg(i, &[b"HINCRBY", key, field, s.as_bytes()]);
Ok(n)
}
pub fn zadd(&mut self, key: &[u8], pairs: &[(f64, &[u8])]) -> KevyResult<usize> {
self.snap(key);
let i = self.idx(key);
let n = self.guards[i]
.store
.zadd(key, pairs)
.map_err(store_err)?;
let score_strs: Vec<Vec<u8>> = pairs
.iter()
.map(|(s, _)| format!("{s}").into_bytes())
.collect();
let mut parts: Vec<&[u8]> = Vec::with_capacity(2 + pairs.len() * 2);
parts.push(b"ZADD");
parts.push(key);
for (j, (_, m)) in pairs.iter().enumerate() {
parts.push(&score_strs[j]);
parts.push(m);
}
self.log_arg(i, &parts);
Ok(n)
}
pub fn zincrby(&mut self, key: &[u8], delta: f64, member: &[u8]) -> KevyResult<f64> {
self.snap(key);
let i = self.idx(key);
let n = self.guards[i]
.store
.zincrby(key, delta, member)
.map_err(store_err)?;
let s = format!("{delta}");
self.log_arg(i, &[b"ZINCRBY", key, s.as_bytes(), member]);
Ok(n)
}
pub fn zscore(&mut self, key: &[u8], member: &[u8]) -> KevyResult<Option<f64>> {
let i = self.idx(key);
self.guards[i].store.zscore(key, member).map_err(store_err)
}
pub fn del(&mut self, keys: &[&[u8]]) -> usize {
for k in keys {
self.snap(k);
}
let mut n = 0;
for k in keys {
let i = self.idx(k);
if self.guards[i].store.del(&[k]) > 0 {
n += 1;
self.log_arg(i, &[b"DEL", k]);
}
}
n
}
pub fn exists(&mut self, keys: &[&[u8]]) -> usize {
keys.iter()
.filter(|k| {
let i = self.idx(k);
self.guards[i].store.key_exists(k)
})
.count()
}
pub fn hdel(&mut self, key: &[u8], fields: &[&[u8]]) -> KevyResult<usize> {
self.snap(key);
let i = self.idx(key);
let removed = self.guards[i].store.hdel(key, fields).map_err(store_err)?;
if removed > 0 {
let mut argv: Vec<&[u8]> = Vec::with_capacity(2 + fields.len());
argv.push(b"HDEL");
argv.push(key);
argv.extend_from_slice(fields);
self.log_arg(i, &argv);
}
Ok(removed)
}
pub fn hgetall(&mut self, key: &[u8]) -> KevyResult<Vec<(Vec<u8>, Vec<u8>)>> {
let i = self.idx(key);
let flat = self.guards[i].store.hgetall(key).map_err(store_err)?;
let mut out = Vec::with_capacity(flat.len() / 2);
let mut it = flat.into_iter();
while let (Some(f), Some(v)) = (it.next(), it.next()) {
out.push((f, v));
}
Ok(out)
}
pub fn hmget(&mut self, key: &[u8], fields: &[&[u8]]) -> KevyResult<Vec<Option<Vec<u8>>>> {
let i = self.idx(key);
self.guards[i].store.hmget(key, fields).map_err(store_err)
}
pub fn hexists(&mut self, key: &[u8], field: &[u8]) -> KevyResult<bool> {
let i = self.idx(key);
self.guards[i].store.hexists(key, field).map_err(store_err)
}
pub fn sadd(&mut self, key: &[u8], members: &[&[u8]]) -> KevyResult<usize> {
self.snap(key);
let i = self.idx(key);
let added = self.guards[i].store.sadd(key, members).map_err(store_err)?;
if added > 0 {
let mut argv: Vec<&[u8]> = Vec::with_capacity(2 + members.len());
argv.push(b"SADD");
argv.push(key);
argv.extend_from_slice(members);
self.log_arg(i, &argv);
}
Ok(added)
}
pub fn srem(&mut self, key: &[u8], members: &[&[u8]]) -> KevyResult<usize> {
self.snap(key);
let i = self.idx(key);
let removed = self.guards[i].store.srem(key, members).map_err(store_err)?;
if removed > 0 {
let mut argv: Vec<&[u8]> = Vec::with_capacity(2 + members.len());
argv.push(b"SREM");
argv.push(key);
argv.extend_from_slice(members);
self.log_arg(i, &argv);
}
Ok(removed)
}
pub fn lpush(&mut self, key: &[u8], values: &[&[u8]]) -> KevyResult<usize> {
self.snap(key);
let i = self.idx(key);
let len = self.guards[i].store.lpush(key, values).map_err(store_err)?;
let mut argv: Vec<&[u8]> = Vec::with_capacity(2 + values.len());
argv.push(b"LPUSH");
argv.push(key);
argv.extend_from_slice(values);
self.log_arg(i, &argv);
Ok(len)
}
pub fn rpush(&mut self, key: &[u8], values: &[&[u8]]) -> KevyResult<usize> {
self.snap(key);
let i = self.idx(key);
let len = self.guards[i].store.rpush(key, values).map_err(store_err)?;
let mut argv: Vec<&[u8]> = Vec::with_capacity(2 + values.len());
argv.push(b"RPUSH");
argv.push(key);
argv.extend_from_slice(values);
self.log_arg(i, &argv);
Ok(len)
}
pub fn zrem(&mut self, key: &[u8], members: &[&[u8]]) -> KevyResult<usize> {
self.snap(key);
let i = self.idx(key);
let removed = self.guards[i].store.zrem(key, members).map_err(store_err)?;
if removed > 0 {
let mut argv: Vec<&[u8]> = Vec::with_capacity(2 + members.len());
argv.push(b"ZREM");
argv.push(key);
argv.extend_from_slice(members);
self.log_arg(i, &argv);
}
Ok(removed)
}
pub fn zcard(&mut self, key: &[u8]) -> KevyResult<usize> {
let i = self.idx(key);
self.guards[i].store.zcard(key).map_err(store_err)
}
pub fn zadd_flags(
&mut self,
key: &[u8],
pairs: &[(f64, &[u8])],
flags: kevy_store::ZaddFlags,
) -> KevyResult<kevy_store::ZaddReport> {
if !flags.valid() {
return Err(KevyError::InvalidInput("invalid ZADD flag combo".into()));
}
let i = self.idx(key);
let rep = self.guards[i]
.store
.zadd_flags(key, pairs, flags)
.map_err(store_err)?;
if !rep.applied.is_empty() {
let score_strs: Vec<Vec<u8>> = rep
.applied
.iter()
.map(|(s, _)| format!("{s}").into_bytes())
.collect();
let mut parts: Vec<&[u8]> = Vec::with_capacity(2 + rep.applied.len() * 2);
parts.push(b"ZADD");
parts.push(key);
for (j, (_, m)) in rep.applied.iter().enumerate() {
parts.push(&score_strs[j]);
parts.push(m);
}
self.log_arg(i, &parts);
}
Ok(rep)
}
}
impl Store {
pub fn atomic_all_shards<R>(
&self,
body: impl FnOnce(&mut AtomicAllShards<'_>) -> KevyResult<R>,
) -> KevyResult<R> {
ensure_writable(self)?;
let guards: Vec<RwLockWriteGuard<'_, Inner>> = self
.shards
.iter()
.map(|s| s.write().expect("lock poisoned"))
.collect();
let mut ctx = AtomicAllShards {
guards,
log: Vec::new(),
undo: Vec::new(),
touched: std::collections::HashSet::new(),
#[cfg(feature = "index")]
indexes: std::sync::Arc::clone(&self.indexes),
};
let outcome = body(&mut ctx);
let log = std::mem::take(&mut ctx.log);
let undo = std::mem::take(&mut ctx.undo);
let r = match outcome {
Ok(r) => r,
Err(e) => {
rollback_all(&mut ctx.guards, undo);
return Err(e);
}
};
commit_group_all(&mut ctx.guards, log)?;
Ok(r)
}
}
#[cfg_attr(not(test), allow(dead_code))]
pub(crate) const ATOMIC_ALL_OPS: &[&str] = &[
"SET", "GET", "INCR", "INCRBY", "HSET", "HGET", "HINCRBY", "ZADD",
"ZINCRBY", "ZSCORE", "DEL", "EXISTS", "HDEL", "HGETALL", "HMGET",
"HEXISTS", "SADD", "SREM", "LPUSH", "RPUSH", "ZREM", "ZCARD",
"SMEMBERS", "SISMEMBER", "LRANGE", "LLEN", "SCARD", "ZRANGEBYSCORE",
];
fn rollback_all(guards: &mut [RwLockWriteGuard<'_, Inner>], undo: Vec<ShardUndoEntry>) {
for (idx, key, prior) in undo.into_iter().rev() {
let g = &mut guards[idx];
match prior {
Some((value, ttl_ms)) => g.store.put_with_ttl(key, value, ttl_ms),
None => {
let k: &[u8] = &key;
g.store.del(&[k]);
}
}
}
}
fn commit_group_all(
guards: &mut [RwLockWriteGuard<'_, Inner>],
log: Vec<(usize, Vec<Vec<u8>>)>,
) -> KevyResult<()> {
#[cfg(feature = "persist")]
for g in guards.iter_mut() {
if let Some(aof) = g.aof.as_mut() {
aof.begin_group();
}
}
let mut commit = Ok(());
for (idx, parts) in log {
let g = &mut guards[idx];
let refs: Vec<&[u8]> = parts.iter().map(|v| v.as_slice()).collect();
commit = commit_write(g, &refs);
if commit.is_err() {
break;
}
}
#[cfg(feature = "persist")]
for g in guards.iter_mut() {
if let Some(aof) = g.aof.as_mut() {
let synced = aof.end_group().map_err(KevyError::from);
if commit.is_ok() {
commit = synced;
}
}
}
commit
}