use std::{sync::Arc, time::Duration};
use bitflags::bitflags;
use super::{
txn_key_entry::{LockType, TxnKeyEntries},
txn_watched_keys_container::TxnWatchedKeysContainer,
watch_version_map::WatchVersionMap,
};
use crate::{
aof::{aof_entry_type::AofEntryType, garnet_log::GarnetLog},
storage::session::storage_session::StoreType,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TxnState {
None,
Started,
Running,
Aborted,
}
bitflags! {
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TransactionStoreTypes: u8 {
const None = 0;
const Main = 1;
const Object = 1 << 1;
const Unified = 1 << 2;
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ClusterSlotVerificationInput {
pub read_only: bool,
pub session_asking: u8,
}
pub trait TxnProcedure {
fn id(&self) -> u8;
fn fail_fast_on_key_lock_failure(&self) -> bool {
false
}
fn key_lock_timeout(&self) -> Duration {
Duration::ZERO
}
fn prepare(&mut self, txn_manager: &mut TransactionManager) -> bool;
fn main(&mut self, txn_manager: &mut TransactionManager, output: &mut Vec<u8>);
fn finalize(&mut self, txn_manager: &mut TransactionManager, output: &mut Vec<u8>);
}
pub struct TransactionGuard<'a> {
txn_manager: Option<&'a mut TransactionManager>,
}
impl TransactionGuard<'_> {
fn null() -> Self {
Self { txn_manager: None }
}
pub fn state(&self) -> TxnState {
self
.txn_manager
.as_deref()
.map_or(TxnState::None, |txn_manager| txn_manager.state)
}
}
impl Drop for TransactionGuard<'_> {
fn drop(&mut self) {
if let Some(txn_manager) = self.txn_manager.as_deref_mut() {
txn_manager.commit(true);
}
}
}
pub struct TransactionManager {
pub state: TxnState,
pub key_entries: TxnKeyEntries,
pub watch_container: TxnWatchedKeysContainer,
pub store_types: TransactionStoreTypes,
pub cluster_enabled: bool,
pub txn_start_head: usize,
pub operation_cnt_txn: usize,
pub perform_writes: bool,
pub txn_version: i64,
pub version_sequence: u64,
pub txn_keys: Vec<Box<[u8]>>,
pub save_key_recv_buffer_ptr: Option<usize>,
pub aof_log: Option<Arc<GarnetLog>>,
pub session_id: i32,
pub stored_proc_mode: bool,
pub is_replaying: bool,
}
impl TransactionManager {
pub fn new(
watch_version_map: Arc<WatchVersionMap>,
aof_log: Option<Arc<GarnetLog>>,
cluster_enabled: bool,
) -> Self {
Self {
state: TxnState::None,
key_entries: TxnKeyEntries::new(16),
watch_container: TxnWatchedKeysContainer::new(watch_version_map),
store_types: TransactionStoreTypes::None,
cluster_enabled,
txn_start_head: 0,
operation_cnt_txn: 0,
perform_writes: false,
txn_version: 0,
version_sequence: 0,
txn_keys: Vec::new(),
save_key_recv_buffer_ptr: None,
aof_log,
session_id: 0,
stored_proc_mode: false,
is_replaying: false,
}
}
pub fn aof_enabled(&self) -> bool {
self.aof_log.is_some()
}
pub fn reset_current(&mut self) {
let is_running = self.state == TxnState::Running;
self.reset(is_running);
}
pub fn reset(&mut self, is_running: bool) {
if is_running {
self.key_entries.unlock_all_keys();
}
self.txn_version = 0;
self.txn_start_head = 0;
self.operation_cnt_txn = 0;
self.state = TxnState::None;
self.store_types = TransactionStoreTypes::None;
self.stored_proc_mode = false;
self.perform_writes = false;
if self.cluster_enabled {
self.txn_keys.clear();
self.save_key_recv_buffer_ptr = None;
}
}
pub fn begin_transaction(&mut self) {}
pub fn locks_acquired(&mut self, _txn_version: i64) {}
pub fn run(
&mut self,
internal_txn: bool,
fail_fast_on_lock: bool,
lock_timeout: Duration,
) -> bool {
if !internal_txn {
let Self {
watch_container,
key_entries,
perform_writes,
..
} = self;
for key in watch_container.save_keys_to_lock() {
Self::register_key_lock(key_entries, perform_writes, key, LockType::Shared);
}
}
self.version_sequence = self.version_sequence.wrapping_add(1);
self.txn_version = self.version_sequence as i64;
self.begin_transaction();
let lock_success = if fail_fast_on_lock {
self.key_entries.try_lock_all_keys(lock_timeout)
} else {
self.key_entries.lock_all_keys();
true
};
if !lock_success || (!internal_txn && !self.watch_container.validate_watch_version()) {
if !lock_success {
log::error!("Transaction failed to acquire all the locks on keys to proceed.");
}
self.reset(true);
if !internal_txn {
self.watch_container.reset();
}
return false;
}
self.locks_acquired(self.txn_version);
if self.perform_writes
&& !self.stored_proc_mode
&& let Some(log) = &self.aof_log
{
let (physical_sublog_access_vector, _) = self.compute_sublog_access_vector();
log.enqueue_txn(
AofEntryType::TxnStart,
self.txn_version,
self.session_id,
&[],
physical_sublog_access_vector,
);
}
self.state = TxnState::Running;
true
}
pub fn commit(&mut self, internal_txn: bool) {
if self.perform_writes
&& !self.stored_proc_mode
&& let Some(log) = &self.aof_log
{
let (physical_sublog_access_vector, _) = self.compute_sublog_access_vector();
log.enqueue_txn(
AofEntryType::TxnCommit,
self.txn_version,
self.session_id,
&[],
physical_sublog_access_vector,
);
}
if !internal_txn {
self.watch_container.reset();
}
self.reset(true);
}
pub fn abort(&mut self) {
self.state = TxnState::Aborted;
}
#[inline]
pub fn is_skipping_operations(&self) -> bool {
self.state == TxnState::Started || self.state == TxnState::Aborted
}
pub fn watch(&mut self, key: &[u8]) {
self.watch_container.add_watch(key);
}
pub fn add_transaction_store_types(&mut self, transaction_store_types: TransactionStoreTypes) {
self.store_types |= transaction_store_types;
}
pub fn add_transaction_store_type(&mut self, store_type: StoreType) {
let transaction_store_types = match store_type {
StoreType::Main => TransactionStoreTypes::Main,
StoreType::Object => TransactionStoreTypes::Object,
StoreType::All => TransactionStoreTypes::Unified,
StoreType::None => TransactionStoreTypes::None,
};
self.store_types |= transaction_store_types;
}
pub fn get_lockset(&self) -> String {
self.key_entries.get_lockset()
}
pub fn get_slot_verification_input(
&mut self,
session_asking: u8,
) -> ClusterSlotVerificationInput {
let Self {
watch_container,
cluster_enabled,
txn_keys,
key_entries,
..
} = self;
for key in watch_container.save_keys_to_key_list() {
if *cluster_enabled {
txn_keys.push(key.into());
}
}
ClusterSlotVerificationInput {
read_only: key_entries.is_read_only(),
session_asking,
}
}
pub fn promote_to_transaction(
&mut self,
store_types: TransactionStoreTypes,
key: &[u8],
lock_type: LockType,
) -> TransactionGuard<'_> {
if self.state == TxnState::Running {
return TransactionGuard::null();
}
self.add_transaction_store_types(store_types);
self.save_key_entry_to_lock(key, lock_type);
let _ = self.run(true, false, Duration::ZERO);
TransactionGuard {
txn_manager: Some(self),
}
}
pub fn compute_custom_proc_sharded_log_access(&self, _key: &[u8]) {}
pub fn compute_sublog_access_vector(&self) -> (u64, u32) {
(0, 0)
}
fn log_proc(&mut self, proc: &dyn TxnProcedure) {
debug_assert!(self.stored_proc_mode);
if self.perform_writes
&& let Some(log) = &self.aof_log
{
let (physical_sublog_access_vector, _) = self.compute_sublog_access_vector();
log.enqueue_stored_proc(
AofEntryType::StoredProcedure,
self.txn_version,
self.session_id,
proc.id(),
&[],
physical_sublog_access_vector,
);
}
}
pub fn run_transaction_proc(
&mut self,
proc: &mut dyn TxnProcedure,
output: &mut Vec<u8>,
is_replaying: bool,
) -> bool {
let running = false;
self.is_replaying = is_replaying;
self.reset_cache_slot_verification_result();
self.stored_proc_mode = true;
if !proc.prepare(self) {
self.reset(running);
return false;
}
if self.state == TxnState::Aborted {
self.write_cached_slot_verification_message(output);
self.reset(running);
return false;
}
if !self.run(
false,
proc.fail_fast_on_key_lock_failure(),
proc.key_lock_timeout(),
) {
self.reset(running);
return false;
}
proc.main(self, output);
if !is_replaying {
self.log_proc(proc);
}
self.commit(false);
if !is_replaying {
proc.finalize(self, output);
}
true
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transaction::{
txn_key_entry_comparison::TxnKeyEntryComparison, watch_version_map::WatchVersionMap,
};
fn manager() -> TransactionManager {
TransactionManager::new(Arc::new(WatchVersionMap::new(64)), None, false)
}
#[test]
fn state_machine_skip_window() {
let mut txn = manager();
assert_eq!(txn.state, TxnState::None);
assert!(!txn.is_skipping_operations());
txn.state = TxnState::Started;
assert!(txn.is_skipping_operations());
txn.state = TxnState::Aborted;
assert!(txn.is_skipping_operations());
txn.state = TxnState::Running;
assert!(!txn.is_skipping_operations());
}
#[test]
fn run_commit_cycle_locks_and_resets() {
let mut txn = manager();
txn.save_key_entry_to_lock(b"alpha", LockType::Exclusive);
assert!(txn.run(false, false, Duration::ZERO));
assert_eq!(txn.state, TxnState::Running);
txn.perform_writes = true;
txn.commit(false);
assert_eq!(txn.state, TxnState::None);
assert_eq!(txn.key_entries.count(), 0);
}
#[test]
fn watch_invalidation_aborts_run() {
let map = Arc::new(WatchVersionMap::new(64));
let mut txn = TransactionManager::new(Arc::clone(&map), None, false);
txn.watch(b"key");
map.increment_version(TxnKeyEntryComparison::key_hash(b"key") as u64);
txn.save_key_entry_to_lock(b"other", LockType::Exclusive);
assert!(!txn.run(false, false, Duration::ZERO));
assert_eq!(txn.state, TxnState::None);
}
#[test]
fn store_type_mapping() {
let mut txn = manager();
txn.add_transaction_store_type(StoreType::Main);
txn.add_transaction_store_type(StoreType::Object);
assert!(txn.store_types.contains(TransactionStoreTypes::Main));
assert!(txn.store_types.contains(TransactionStoreTypes::Object));
assert!(!txn.store_types.contains(TransactionStoreTypes::Unified));
txn.add_transaction_store_type(StoreType::All);
assert!(txn.store_types.contains(TransactionStoreTypes::Unified));
}
#[test]
fn guard_commits_on_drop() {
let mut txn = manager();
{
let guard =
txn.promote_to_transaction(TransactionStoreTypes::Main, b"k", LockType::Exclusive);
assert_eq!(guard.state(), TxnState::Running);
}
assert_eq!(txn.state, TxnState::None);
}
#[test]
fn guard_is_null_when_already_running() {
let mut txn = manager();
txn.save_key_entry_to_lock(b"k", LockType::Shared);
assert!(txn.run(false, false, Duration::ZERO));
{
let guard = txn.promote_to_transaction(TransactionStoreTypes::Main, b"j", LockType::Shared);
assert_eq!(guard.state(), TxnState::None);
}
assert_eq!(txn.state, TxnState::Running);
}
struct CountingProc {
prepared: bool,
finalized: bool,
}
impl TxnProcedure for CountingProc {
fn id(&self) -> u8 {
9
}
fn prepare(&mut self, txn: &mut TransactionManager) -> bool {
txn.save_key_entry_to_lock(b"proc-key", LockType::Exclusive);
self.prepared = true;
true
}
fn main(&mut self, _txn: &mut TransactionManager, output: &mut Vec<u8>) {
output.extend_from_slice(b"main");
}
fn finalize(&mut self, _txn: &mut TransactionManager, _output: &mut Vec<u8>) {
self.finalized = true;
}
}
#[test]
fn run_transaction_proc_full_flow() {
let mut txn = manager();
let mut proc = CountingProc {
prepared: false,
finalized: false,
};
let mut output = Vec::new();
assert!(txn.run_transaction_proc(&mut proc, &mut output, false));
assert!(proc.prepared && proc.finalized);
assert_eq!(output, b"main");
assert_eq!(txn.state, TxnState::None);
}
#[test]
fn run_transaction_proc_skips_finalize_on_replay() {
let mut txn = manager();
let mut proc = CountingProc {
prepared: false,
finalized: false,
};
let mut output = Vec::new();
assert!(txn.run_transaction_proc(&mut proc, &mut output, true));
assert!(proc.prepared);
assert!(!proc.finalized);
}
}