/* * SPDX-FileCopyrightText: 2020 Stalwart Labs LLC * * SPDX-License-Identifier: AGPL-3.0-only OR LicenseRef-SEL * * Modified by Coffey Labs in 2026 for INBUXA. */ use registry::schema::structs::Rate; use std::borrow::Cow; use trc::AddContext; #[allow(unused_imports)] use crate::{ Deserialize, InMemoryStore, IterateParams, QueryResult, Store, U64_LEN, Value, ValueKey, write::{ BatchBuilder, Operation, ValueClass, ValueOp, key::{DeserializeBigEndian, KeySerializer}, now, }, }; use crate::{ SerializeInfallible, backend::{http::lookup::HttpStoreGet, memory::StaticMemoryStore}, write::{InMemoryClass, assert::AssertValue}, }; pub struct KeyValue { pub key: Vec, pub value: T, pub expires: Option, } impl InMemoryStore { pub async fn key_set(&self, kv: KeyValue>) -> trc::Result<()> { match self { InMemoryStore::Store(store) => { let mut batch = BatchBuilder::new(); batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Key(kv.key)), op: ValueOp::Set( KeySerializer::new(kv.value.len() + U64_LEN) .write(kv.expires.map_or(u64::MAX, |expires| now() + expires)) .write(kv.value.as_slice()) .finalize(), ), }); store.write(batch.build_all()).await.map(|_| ()) } // inbuxa: ST-23 InMemoryStore::Sharded(store) => { let member = store.member(&kv.key); Box::pin(member.key_set(kv)).await } #[cfg(feature = "redis")] InMemoryStore::Redis(store) => store.key_set(&kv.key, &kv.value, kv.expires).await, InMemoryStore::Static(_) | InMemoryStore::Http(_) => { Err(trc::StoreEvent::NotSupported.into_err()) } } .caused_by(trc::location!()) } pub async fn counter_incr(&self, kv: KeyValue, return_value: bool) -> trc::Result { match self { InMemoryStore::Store(store) => { let mut batch = BatchBuilder::new(); if let Some(expires) = kv.expires { batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Key(kv.key.clone())), op: ValueOp::Set( KeySerializer::new(U64_LEN * 2) .write(0u64) .write(now() + expires) .finalize(), ), }); } if return_value { batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Counter(kv.key)), op: ValueOp::AddAndGet(kv.value), }); store .write(batch.build_all()) .await .and_then(|r| r.last_counter_id()) } else { batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Counter(kv.key)), op: ValueOp::AtomicAdd(kv.value), }); store.write(batch.build_all()).await.map(|_| 0) } } // inbuxa: ST-23 InMemoryStore::Sharded(store) => { let member = store.member(&kv.key); Box::pin(member.counter_incr(kv, return_value)).await } #[cfg(feature = "redis")] InMemoryStore::Redis(store) => store.key_incr(&kv.key, kv.value, kv.expires).await, InMemoryStore::Static(_) | InMemoryStore::Http(_) => { Err(trc::StoreEvent::NotSupported.into_err()) } } .caused_by(trc::location!()) } pub async fn key_delete(&self, key: impl Into>) -> trc::Result<()> { match self { InMemoryStore::Store(store) => { let mut batch = BatchBuilder::new(); batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Key(key.into().into_bytes())), op: ValueOp::Clear, }); store.write(batch.build_all()).await.map(|_| ()) } // inbuxa: ST-23 InMemoryStore::Sharded(store) => { let key = key.into().into_bytes(); Box::pin(store.member(&key).key_delete(key.clone())).await } #[cfg(feature = "redis")] InMemoryStore::Redis(store) => store.key_delete(key.into().as_bytes()).await, InMemoryStore::Static(_) | InMemoryStore::Http(_) => { Err(trc::StoreEvent::NotSupported.into_err()) } } .caused_by(trc::location!()) } pub async fn counter_delete(&self, key: impl Into>) -> trc::Result<()> { match self { InMemoryStore::Store(store) => { let mut batch = BatchBuilder::new(); batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Counter(key.into().into_bytes())), op: ValueOp::Clear, }); store.write(batch.build_all()).await.map(|_| ()) } // inbuxa: ST-23 InMemoryStore::Sharded(store) => { let key = key.into().into_bytes(); Box::pin(store.member(&key).counter_delete(key.clone())).await } #[cfg(feature = "redis")] InMemoryStore::Redis(store) => store.key_delete(key.into().as_bytes()).await, InMemoryStore::Static(_) | InMemoryStore::Http(_) => { Err(trc::StoreEvent::NotSupported.into_err()) } } .caused_by(trc::location!()) } pub async fn key_delete_prefix(&self, prefix: &[u8]) -> trc::Result<()> { match self { InMemoryStore::Store(store) => { if prefix.is_empty() { return Ok(()); } let from_range = prefix.to_vec(); let mut to_range = Vec::with_capacity(prefix.len() + 3); to_range.extend_from_slice(prefix); to_range.extend_from_slice([u8::MAX, u8::MAX, u8::MAX].as_ref()); store .delete_range( ValueKey::from(ValueClass::InMemory(InMemoryClass::Counter( from_range.clone(), ))), ValueKey::from(ValueClass::InMemory(InMemoryClass::Counter( to_range.clone(), ))), ) .await?; store .delete_range( ValueKey::from(ValueClass::InMemory(InMemoryClass::Key(from_range))), ValueKey::from(ValueClass::InMemory(InMemoryClass::Key(to_range))), ) .await } // inbuxa: ST-24: every member InMemoryStore::Sharded(store) => { for (index, member) in store.members.iter().enumerate() { Box::pin(member.key_delete_prefix(prefix)) .await .map_err(|err| err.details(format!("Member {}", index + 1)))?; } Ok(()) } #[cfg(feature = "redis")] InMemoryStore::Redis(store) => store.key_delete_prefix(prefix).await, InMemoryStore::Static(_) | InMemoryStore::Http(_) => { Err(trc::StoreEvent::NotSupported.into_err()) } } .caused_by(trc::location!()) } pub async fn key_get> + std::fmt::Debug + 'static>( &self, key: impl Into>, ) -> trc::Result> { match self { InMemoryStore::Store(store) => store .get_value::>(ValueKey::from(ValueClass::InMemory( InMemoryClass::Key(key.into().into_bytes()), ))) .await .map(|value| value.and_then(|v| v.into())), // inbuxa: ST-23 InMemoryStore::Sharded(store) => { let key = key.into().into_bytes(); Box::pin(store.member(&key).key_get::(key.clone())).await } #[cfg(feature = "redis")] InMemoryStore::Redis(store) => store.key_get(key.into().as_bytes()).await, InMemoryStore::Static(store) => Ok(match store.as_ref() { StaticMemoryStore::Map(map) => map .get(key.into().as_str()) .map(|value| T::from(value.clone())), StaticMemoryStore::Set(set) => { if set.contains(key.into().as_str()) { Some(T::from(Value::Bool(true))) } else { None } } }), InMemoryStore::Http(store) => { Ok(store.get(key.into().as_str()).map(|value| T::from(value))) } } .caused_by(trc::location!()) } pub async fn counter_get(&self, key: impl Into>) -> trc::Result { match self { InMemoryStore::Store(store) => { store .get_counter(ValueKey::from(ValueClass::InMemory( InMemoryClass::Counter(key.into().into_bytes()), ))) .await } // inbuxa: ST-23 InMemoryStore::Sharded(store) => { let key = key.into().into_bytes(); Box::pin(store.member(&key).counter_get(key.clone())).await } #[cfg(feature = "redis")] InMemoryStore::Redis(store) => store.counter_get(key.into().as_bytes()).await, InMemoryStore::Static(_) | InMemoryStore::Http(_) => { Err(trc::StoreEvent::NotSupported.into_err()) } } .caused_by(trc::location!()) } pub async fn key_exists(&self, key: impl Into>) -> trc::Result { match self { InMemoryStore::Store(store) => store .get_value::>(ValueKey::from(ValueClass::InMemory( InMemoryClass::Key(key.into().into_bytes()), ))) .await .map(|value| matches!(value, Some(LookupValue::Value(Empty)))), // inbuxa: ST-23 InMemoryStore::Sharded(store) => { let key = key.into().into_bytes(); Box::pin(store.member(&key).key_exists(key.clone())).await } #[cfg(feature = "redis")] InMemoryStore::Redis(store) => store.key_exists(key.into().as_bytes()).await, InMemoryStore::Static(store) => Ok(match store.as_ref() { StaticMemoryStore::Map(map) => map.get(key.into().as_str()).is_some(), StaticMemoryStore::Set(set) => set.contains(key.into().as_str()), }), InMemoryStore::Http(store) => Ok(store.contains(key.into().as_str())), } .caused_by(trc::location!()) } pub async fn is_rate_allowed( &self, prefix: u8, key: &[u8], rate: &Rate, soft_check: bool, ) -> trc::Result> { let now = now(); let period = rate.period.as_secs().max(1); let range_start = now / period; let range_end = (range_start * period) + period; let expires_in = range_end - now; let mut bucket = Vec::with_capacity(key.len() + U64_LEN + 1); bucket.push(prefix); bucket.extend_from_slice(key); bucket.extend_from_slice(range_start.to_be_bytes().as_slice()); let requests = if !soft_check { self.counter_incr(KeyValue::new(bucket, 1).expires(expires_in), true) .await .caused_by(trc::location!())? } else { self.counter_get(bucket).await.caused_by(trc::location!())? + 1 }; if requests <= rate.count as i64 { Ok(None) } else { Ok(Some(expires_in)) } } pub async fn try_lock(&self, prefix: u8, key: &[u8], duration: u64) -> trc::Result { match self { InMemoryStore::Store(store) => { let key = KeyValue::<()>::build_key(prefix, key); let lock_expiry = match store .get_value::(ValueKey::from(ValueClass::InMemory(InMemoryClass::Key( key.clone(), )))) .await { Ok(lock_expiry) => lock_expiry, Err(err) if err.matches(trc::EventType::Store(trc::StoreEvent::DataCorruption)) => { // TODO remove in 1.0 let mut batch = BatchBuilder::new(); batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Key(key.clone())), op: ValueOp::Clear, }); store .write(batch.build_all()) .await .caused_by(trc::location!())?; None } Err(err) => { return Err(err .details("Failed to read lock.") .caused_by(trc::location!())); } }; let now = now(); if lock_expiry.is_some_and(|expiry| expiry > now) { return Ok(false); } let key: ValueClass = ValueClass::InMemory(InMemoryClass::Key(key)); let mut batch = BatchBuilder::new(); batch.assert_value( key.clone(), match lock_expiry { Some(value) => AssertValue::U64(value), None => AssertValue::None, }, ); batch.set(key.clone(), (now + duration).serialize()); match store.write(batch.build_all()).await { Ok(_) => Ok(true), Err(err) if err.is_assertion_failure() => Ok(false), Err(err) => Err(err .details("Failed to lock event.") .caused_by(trc::location!())), } } // inbuxa: ST-23: a lock and its release meet on one member InMemoryStore::Sharded(store) => { Box::pin( store .member(&KeyValue::<()>::build_key(prefix, key)) .try_lock(prefix, key, duration), ) .await } #[cfg(feature = "redis")] InMemoryStore::Redis(store) => { store .try_lock(&KeyValue::<()>::build_key(prefix, key), duration) .await } InMemoryStore::Static(_) | InMemoryStore::Http(_) => { Err(trc::StoreEvent::NotSupported.into_err()) } } } pub async fn remove_lock(&self, prefix: u8, key: &[u8]) -> trc::Result<()> { self.key_delete(KeyValue::<()>::build_key(prefix, key)) .await } pub async fn purge_in_memory_store(&self) -> trc::Result<()> { match self { InMemoryStore::Store(store) => { // Delete expired keys and counters let from_key = ValueKey::from(ValueClass::InMemory(InMemoryClass::Key(vec![0u8]))); let to_key = ValueKey::from(ValueClass::InMemory(InMemoryClass::Key(vec![u8::MAX; 10]))); let current_time = now(); let mut expired_keys = Vec::new(); let mut expired_counters = Vec::new(); store .iterate(IterateParams::new(from_key, to_key), |key, value| { let expiry = value.deserialize_be_u64(0).caused_by(trc::location!())?; if expiry == 0 { if value .deserialize_be_u64(U64_LEN) .caused_by(trc::location!())? <= current_time { expired_counters.push(key.to_vec()); } } else if expiry <= current_time { expired_keys.push(key.to_vec()); } Ok(true) }) .await .caused_by(trc::location!())?; if !expired_keys.is_empty() { let mut batch = BatchBuilder::new(); for key in expired_keys { batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Key(key)), op: ValueOp::Clear, }); if batch.is_large_batch() { store .write(batch.build_all()) .await .caused_by(trc::location!())?; batch = BatchBuilder::new(); } } if !batch.is_empty() { store .write(batch.build_all()) .await .caused_by(trc::location!())?; } } if !expired_counters.is_empty() { let mut batch = BatchBuilder::new(); for key in expired_counters { batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Counter(key.clone())), op: ValueOp::Clear, }); batch.any_op(Operation::Value { class: ValueClass::InMemory(InMemoryClass::Key(key)), op: ValueOp::Clear, }); if batch.is_large_batch() { store .write(batch.build_all()) .await .caused_by(trc::location!())?; batch = BatchBuilder::new(); } } if !batch.is_empty() { store .write(batch.build_all()) .await .caused_by(trc::location!())?; } } } // inbuxa: ST-24: every member InMemoryStore::Sharded(store) => { for (index, member) in store.members.iter().enumerate() { Box::pin(member.purge_in_memory_store()) .await .map_err(|err| err.details(format!("Member {}", index + 1)))?; } } #[cfg(feature = "redis")] InMemoryStore::Redis(_) => {} InMemoryStore::Static(_) | InMemoryStore::Http(_) => {} } Ok(()) } pub fn is_sql(&self) -> bool { match self { InMemoryStore::Store(store) => store.is_sql(), _ => false, } } pub fn is_redis(&self) -> bool { match self { #[cfg(feature = "redis")] InMemoryStore::Redis(_) => true, // inbuxa: ST-3: as its members are InMemoryStore::Sharded(store) => store.members.iter().all(|m| m.is_redis()), InMemoryStore::Static(_) => false, _ => false, } } pub fn into_store(self) -> Option { match self { InMemoryStore::Store(store) => Some(store), _ => None, } } } pub enum LookupKey<'x> { String(String), StringRef(&'x str), Bytes(Vec), BytesRef(&'x [u8]), } impl<'x> From<&'x str> for LookupKey<'x> { fn from(key: &'x str) -> Self { LookupKey::StringRef(key) } } impl<'x> From<&'x String> for LookupKey<'x> { fn from(key: &'x String) -> Self { LookupKey::StringRef(key.as_str()) } } impl<'x> From<&'x [u8]> for LookupKey<'x> { fn from(key: &'x [u8]) -> Self { LookupKey::BytesRef(key) } } impl<'x> From> for LookupKey<'x> { fn from(key: Cow<'x, str>) -> Self { match key { Cow::Borrowed(key) => LookupKey::StringRef(key), Cow::Owned(key) => LookupKey::String(key), } } } impl From for LookupKey<'static> { fn from(key: String) -> Self { LookupKey::String(key) } } impl From> for LookupKey<'static> { fn from(key: Vec) -> Self { LookupKey::Bytes(key) } } impl LookupKey<'_> { pub fn as_str(&self) -> &str { match self { LookupKey::String(string) => string, LookupKey::StringRef(string) => string, LookupKey::Bytes(bytes) => std::str::from_utf8(bytes).unwrap_or_default(), LookupKey::BytesRef(bytes) => std::str::from_utf8(bytes).unwrap_or_default(), } } pub fn into_bytes(self) -> Vec { match self { LookupKey::String(string) => string.into_bytes(), LookupKey::StringRef(string) => string.as_bytes().to_vec(), LookupKey::Bytes(bytes) => bytes, LookupKey::BytesRef(bytes) => bytes.to_vec(), } } pub fn as_bytes(&self) -> &[u8] { match self { LookupKey::String(string) => string.as_bytes(), LookupKey::StringRef(string) => string.as_bytes(), LookupKey::Bytes(bytes) => bytes.as_slice(), LookupKey::BytesRef(bytes) => bytes, } } } impl KeyValue { pub fn build_key(prefix: u8, key: impl AsRef<[u8]>) -> Vec { let key_ = key.as_ref(); let mut key = Vec::with_capacity(key_.len() + 1); key.push(prefix); key.extend_from_slice(key_); key } pub fn with_prefix(prefix: u8, key: impl AsRef<[u8]>, value: T) -> Self { Self { key: Self::build_key(prefix, key), value, expires: None, } } pub fn new(key: impl Into>, value: T) -> Self { Self { key: key.into(), value, expires: None, } } pub fn expires(mut self, expires: u64) -> Self { self.expires = expires.into(); self } pub fn expires_opt(mut self, expires: Option) -> Self { self.expires = expires; self } } struct Empty; enum LookupValue { Value(T), None, } impl Deserialize for LookupValue { fn deserialize(bytes: &[u8]) -> trc::Result { bytes.deserialize_be_u64(0).and_then(|expires| { Ok(if expires > now() { LookupValue::Value( T::deserialize(bytes.get(U64_LEN..).unwrap_or_default()) .caused_by(trc::location!())?, ) } else { LookupValue::None }) }) } } impl Deserialize for Empty { fn deserialize(_bytes: &[u8]) -> trc::Result { Ok(Empty) } } impl From> for Option { fn from(value: LookupValue) -> Self { match value { LookupValue::Value(value) => Some(value), LookupValue::None => None, } } } impl From> for String { fn from(value: Value<'static>) -> Self { match value { Value::Text(string) => string.into_owned(), Value::Blob(bytes) => String::from_utf8_lossy(bytes.as_ref()).into_owned(), Value::Bool(boolean) => boolean.to_string(), Value::Null => String::new(), Value::Integer(num) => num.to_string(), Value::Float(num) => num.to_string(), } } }