use std::time::Duration;
use super::{
transaction_manager::{TransactionManager, TxnState},
txn_key_manager::TxnCommandKeys,
};
use crate::{
objects::parse_utils::try_get_int,
resp::{
cmd_strings::{
GENERIC_ERR_COMMAND_DISALLOWED_WITH_OPTION, GENERIC_ERR_WRONG_NUM_ARGS,
RESP_ERR_GENERIC_UNK_CMD, RESP_ERR_GENERIC_VALUE_IS_NOT_INTEGER, RESP_OK,
},
resp_server_session::RespServerSession,
},
storage::session::storage_session::StoreType,
types::RespCommand,
};
const RESP_QUEUED: &[u8] = b"+QUEUED\r\n";
const RESP_ERR_GENERIC_NESTED_MULTI: &str = "ERR MULTI calls can not be nested";
const RESP_ERR_EXEC_ABORT: &str = "EXECABORT Transaction discarded because of previous errors.";
const RESP_ERR_GENERIC_EXEC_WO_MULTI: &str = "ERR EXEC without MULTI";
const RESP_ERR_GENERIC_DISCARD_WO_MULTI: &str = "ERR DISCARD without MULTI";
const RESP_ERR_GENERIC_WATCH_IN_MULTI: &str = "ERR WATCH inside MULTI is not allowed";
const RESP_ERR_SELECT_IN_TXN_UNSUPPORTED: &str =
"ERR SELECT is currently unsupported inside a transaction.";
const RESP_ERR_SWAPDB_IN_TXN_UNSUPPORTED: &str =
"ERR SWAPDB is currently unsupported inside a transaction.";
const RESP_ERR_NO_TRANSACTION_PROCEDURE: &str = "ERR Could not get transaction procedure";
const GENERIC_ERR_WRONG_NUM_ARGS_TXN: &str =
"ERR Invalid number of parameters to stored proc {0}, expected {1}, actual {2}";
#[derive(Debug, Clone)]
pub struct TxnQueuedCommandInfo {
pub name: String,
pub arity: i32,
pub allowed_in_txn: bool,
pub is_sub_command: bool,
pub keys: Option<TxnCommandKeys>,
}
pub struct TxnProcHandle {
pub name: String,
pub arity: i32,
}
pub trait TxnProcResolver {
fn get_custom_transaction_procedure(&self, txn_id: u8) -> Option<TxnProcHandle>;
fn try_transaction_proc(
&mut self,
txn_id: u8,
txn_manager: &mut TransactionManager,
session: &mut RespServerSession,
) -> bool;
}
fn write_error(session: &mut RespServerSession, message: &str) {
session.abort_error_message(message);
}
fn write_array_length(session: &mut RespServerSession, count: usize) {
session
.output
.extend_from_slice(format!("*{count}\r\n").as_bytes());
}
fn write_null_array(session: &mut RespServerSession) {
session.output.extend_from_slice(b"*-1\r\n");
}
impl TransactionManager {
pub fn network_multi(&mut self, session: &mut RespServerSession) -> bool {
if self.state != TxnState::None {
write_error(session, RESP_ERR_GENERIC_NESTED_MULTI);
self.abort();
return true;
}
self.txn_start_head = session.end_read_head;
self.state = TxnState::Started;
self.operation_cnt_txn = 0;
self.save_key_recv_buffer_ptr = None;
session.output.extend_from_slice(RESP_OK);
true
}
pub fn network_exec(&mut self, session: &mut RespServerSession) -> bool {
if self.state == TxnState::Running {
self.commit(false);
return true;
}
if self.state == TxnState::Aborted {
write_error(session, RESP_ERR_EXEC_ABORT);
self.reset(false);
self.watch_container.reset();
return true;
}
if self.state == TxnState::Started {
let orig_read_head = session.end_read_head;
session.end_read_head = self.txn_start_head;
if self.cluster_enabled {
let _ = self.get_slot_verification_input(session.session_asking);
}
let start_txn = self.run(false, false, Duration::ZERO);
if start_txn {
write_array_length(session, self.operation_cnt_txn);
} else {
session.end_read_head = orig_read_head;
write_null_array(session);
}
return true;
}
write_error(session, RESP_ERR_GENERIC_EXEC_WO_MULTI);
true
}
pub fn network_skip(
&mut self,
session: &mut RespServerSession,
cmd: RespCommand,
info: Option<&TxnQueuedCommandInfo>,
) -> bool {
let Some(command_info) = info.filter(|info| info.allowed_in_txn) else {
write_error(session, RESP_ERR_GENERIC_UNK_CMD);
self.abort();
return true;
};
let count = session.parse_state.count;
let mut arity = if command_info.arity > 0 {
command_info.arity - 1
} else {
command_info.arity + 1
};
if command_info.is_sub_command || cmd == RespCommand::Bitop {
arity = if arity > 0 { arity - 1 } else { arity + 1 };
}
let invalid_num_args = if arity > 0 {
count != arity as usize
} else {
(count as i64) < -(arity as i64)
};
let is_watch = matches!(
cmd,
RespCommand::Watch | RespCommand::Watchms | RespCommand::Watchos
);
let is_multi_db_command = matches!(cmd, RespCommand::Select | RespCommand::Swapdb);
if invalid_num_args || is_watch || is_multi_db_command {
if is_watch {
write_error(session, RESP_ERR_GENERIC_WATCH_IN_MULTI);
return true;
}
if invalid_num_args {
write_error(
session,
&GENERIC_ERR_WRONG_NUM_ARGS.replace("{0}", &command_info.name),
);
self.abort();
return true;
}
match cmd {
RespCommand::Swapdb => {
write_error(session, RESP_ERR_SWAPDB_IN_TXN_UNSUPPORTED);
self.abort();
return true;
}
RespCommand::Select => {
if count > 0
&& let Some(index) = try_get_int(session.parse_state.get_arg_slice_by_ref(0).as_slice())
&& i64::from(index) != session.active_db_id
{
write_error(session, RESP_ERR_SELECT_IN_TXN_UNSUPPORTED);
self.abort();
return true;
}
}
_ => {}
}
}
if cmd == RespCommand::Debug && !session.can_run_debug() {
let message = GENERIC_ERR_COMMAND_DISALLOWED_WITH_OPTION
.replace("{0}", "DEBUG")
.replace("{1}", "enable-debug-command");
write_error(session, &message);
self.abort();
return true;
}
if self.cluster_enabled {
self.copy_existing_keys_to_scratch_buffer();
}
if let Some(keys) = &command_info.keys {
self.lock_keys(session, keys);
}
session.output.extend_from_slice(RESP_QUEUED);
self.operation_cnt_txn += 1;
true
}
pub fn network_discard(&mut self, session: &mut RespServerSession) -> bool {
if self.state == TxnState::None {
write_error(session, RESP_ERR_GENERIC_DISCARD_WO_MULTI);
return true;
}
session.output.extend_from_slice(RESP_OK);
self.reset(false);
self.watch_container.reset();
true
}
pub fn common_watch(&mut self, session: &mut RespServerSession, store_type: StoreType) -> bool {
let count = session.parse_state.count;
if count == 0 {
write_error(session, GENERIC_ERR_WRONG_NUM_ARGS);
return true;
}
self.add_transaction_store_type(store_type);
for c in 0..count {
let key = session.parse_state.get_arg_slice_by_ref(c);
self.watch(key.as_slice());
}
session.output.extend_from_slice(RESP_OK);
true
}
pub fn network_watch_ms(&mut self, session: &mut RespServerSession) -> bool {
self.common_watch(session, StoreType::Main)
}
pub fn network_watch_os(&mut self, session: &mut RespServerSession) -> bool {
self.common_watch(session, StoreType::Object)
}
pub fn network_watch(&mut self, session: &mut RespServerSession) -> bool {
self.common_watch(session, StoreType::All)
}
pub fn network_unwatch(&mut self, session: &mut RespServerSession) -> bool {
if self.state == TxnState::None {
self.watch_container.reset();
}
session.output.extend_from_slice(RESP_OK);
true
}
pub fn network_runtxp_fast(
&mut self,
session: &mut RespServerSession,
resolver: &mut dyn TxnProcResolver,
) -> bool {
self.network_runtxp(session, resolver)
}
pub fn network_runtxp(
&mut self,
session: &mut RespServerSession,
resolver: &mut dyn TxnProcResolver,
) -> bool {
let count = session.parse_state.count;
if count < 1 {
session.abort_wrong_num_args("RUNTXP");
return true;
}
let first = session.parse_state.get_arg_slice_by_ref(0);
let Some(tx_id) = try_get_int(first.as_slice()) else {
write_error(session, RESP_ERR_GENERIC_VALUE_IS_NOT_INTEGER);
return true;
};
let Ok(tx_id) = u8::try_from(tx_id) else {
write_error(session, RESP_ERR_NO_TRANSACTION_PROCEDURE);
return true;
};
let Some(proc) = resolver.get_custom_transaction_procedure(tx_id) else {
write_error(session, RESP_ERR_NO_TRANSACTION_PROCEDURE);
return true;
};
if (proc.arity > 0 && count != proc.arity as usize)
|| (proc.arity < 0 && (count as i64) < -(proc.arity as i64))
{
let expected_params = if proc.arity > 0 {
proc.arity - 1
} else {
-proc.arity - 1
};
write_error(
session,
&GENERIC_ERR_WRONG_NUM_ARGS_TXN
.replace("{0}", &tx_id.to_string())
.replace("{1}", &expected_params.to_string())
.replace("{2}", &(count - 1).to_string()),
);
return true;
}
resolver.try_transaction_proc(tx_id, self, session);
true
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::{
super::{
transaction_manager::TransactionManager, txn_key_manager::TxnKeySpec,
watch_version_map::WatchVersionMap,
},
*,
};
use crate::{arg_slice::ArgSlice, resp::resp_server_session::RespServerSession};
fn manager() -> TransactionManager {
TransactionManager::new(Arc::new(WatchVersionMap::new(64)), None, false)
}
fn session_with_args(args: &[&[u8]]) -> (RespServerSession, Vec<u8>) {
let mut buffer: Vec<u8> = Vec::new();
for arg in args {
buffer.extend_from_slice(arg);
}
let mut slices = Vec::with_capacity(args.len());
let mut offset = 0usize;
for arg in args {
slices.push(ArgSlice::new(
unsafe { buffer.as_ptr().add(offset) },
arg.len(),
));
offset += arg.len();
}
let mut session = RespServerSession::default();
session.parse_state.initialize_with_args(&slices);
(session, buffer)
}
#[test]
fn multi_then_nested_multi_aborts() {
let mut txn = manager();
let mut session = RespServerSession::default();
assert!(txn.network_multi(&mut session));
assert_eq!(txn.state, TxnState::Started);
assert_eq!(session.output, b"+OK\r\n");
session.output.clear();
assert!(txn.network_multi(&mut session));
assert_eq!(txn.state, TxnState::Aborted);
assert_eq!(
session.output,
format!("-{RESP_ERR_GENERIC_NESTED_MULTI}\r\n").as_bytes()
);
}
#[test]
fn exec_without_multi_errors() {
let mut txn = manager();
let mut session = RespServerSession::default();
assert!(txn.network_exec(&mut session));
assert_eq!(
session.output,
format!("-{RESP_ERR_GENERIC_EXEC_WO_MULTI}\r\n").as_bytes()
);
}
#[test]
fn discard_without_multi_errors() {
let mut txn = manager();
let mut session = RespServerSession::default();
assert!(txn.network_discard(&mut session));
assert_eq!(
session.output,
format!("-{RESP_ERR_GENERIC_DISCARD_WO_MULTI}\r\n").as_bytes()
);
assert_eq!(txn.state, TxnState::None);
session.output.clear();
assert!(txn.network_multi(&mut session));
session.output.clear();
assert!(txn.network_discard(&mut session));
assert_eq!(session.output, b"+OK\r\n");
assert_eq!(txn.state, TxnState::None);
}
#[test]
fn skip_queues_and_counts() {
let mut txn = manager();
let (mut session, _buffer) = session_with_args(&[b"k1", b"v1"]);
let info = TxnQueuedCommandInfo {
name: "set".into(),
arity: 3, allowed_in_txn: true,
is_sub_command: false,
keys: Some(TxnCommandKeys {
store_type: StoreType::Main,
key_specs: vec![TxnKeySpec::new(0, 0, 1, false)],
}),
};
assert!(txn.network_skip(&mut session, RespCommand::Set, Some(&info)));
assert_eq!(session.output, RESP_QUEUED);
assert_eq!(txn.operation_cnt_txn, 1);
assert_eq!(txn.key_entries.count(), 1);
assert!(txn.perform_writes);
}
#[test]
fn skip_unknown_command_aborts() {
let mut txn = manager();
let (mut session, _buffer) = session_with_args(&[]);
assert!(txn.network_skip(&mut session, RespCommand::Invalid, None));
assert_eq!(txn.state, TxnState::Aborted);
assert!(String::from_utf8_lossy(&session.output).contains("unknown command"));
}
#[test]
fn skip_watch_inside_multi_errors_without_abort() {
let mut txn = manager();
let (mut session, _buffer) = session_with_args(&[b"k"]);
txn.state = TxnState::Started;
let watch_info = TxnQueuedCommandInfo {
name: "watch".into(),
arity: -2,
allowed_in_txn: true,
is_sub_command: false,
keys: None,
};
assert!(txn.network_skip(&mut session, RespCommand::Watch, Some(&watch_info)));
assert_eq!(txn.state, TxnState::Started); assert_eq!(txn.operation_cnt_txn, 0); assert!(String::from_utf8_lossy(&session.output).contains("WATCH inside MULTI is not allowed"));
}
#[test]
fn skip_wrong_arity_aborts() {
let mut txn = manager();
let (mut session, _buffer) = session_with_args(&[b"only-key"]);
let info = TxnQueuedCommandInfo {
name: "rpush".into(),
arity: -3, allowed_in_txn: true,
is_sub_command: false,
keys: None,
};
assert!(txn.network_skip(&mut session, RespCommand::Rpush, Some(&info)));
assert_eq!(txn.state, TxnState::Aborted);
assert!(String::from_utf8_lossy(&session.output).contains("wrong number of arguments"));
}
#[test]
fn swapdb_inside_multi_aborts() {
let mut txn = manager();
let (mut session, _buffer) = session_with_args(&[b"0", b"1"]);
let info = TxnQueuedCommandInfo {
name: "swapdb".into(),
arity: 3,
allowed_in_txn: true,
is_sub_command: false,
keys: None,
};
assert!(txn.network_skip(&mut session, RespCommand::Swapdb, Some(&info)));
assert_eq!(txn.state, TxnState::Aborted);
assert!(String::from_utf8_lossy(&session.output).contains("SWAPDB is currently unsupported"));
}
#[test]
fn unwatch_resets_watches_when_idle() {
let mut txn = manager();
txn.watch(b"k");
let mut session = RespServerSession::default();
assert!(txn.network_unwatch(&mut session));
assert_eq!(session.output, b"+OK\r\n");
assert!(txn.watch_container.validate_watch_version());
}
struct MockResolver;
impl TxnProcResolver for MockResolver {
fn get_custom_transaction_procedure(&self, txn_id: u8) -> Option<TxnProcHandle> {
(txn_id == 7).then(|| TxnProcHandle {
name: "mock-proc".into(),
arity: 2,
})
}
fn try_transaction_proc(
&mut self,
_txn_id: u8,
_txn_manager: &mut TransactionManager,
session: &mut RespServerSession,
) -> bool {
session.output.extend_from_slice(b"PROC-MAIN");
true
}
}
#[test]
fn runtxp_routes_through_resolver() {
let mut txn = manager();
let (mut session, _buffer) = session_with_args(&[b"7", b"a1"]);
assert!(txn.network_runtxp(&mut session, &mut MockResolver));
assert_eq!(session.output, b"PROC-MAIN");
}
#[test]
fn runtxp_unknown_procedure_errors() {
let mut txn = manager();
let (mut session, _buffer) = session_with_args(&[b"99"]);
assert!(txn.network_runtxp(&mut session, &mut MockResolver));
assert!(
String::from_utf8_lossy(&session.output).contains("Could not get transaction procedure")
);
}
#[test]
fn runtxp_non_integer_id_errors() {
let mut txn = manager();
let (mut session, _buffer) = session_with_args(&[b"abc"]);
assert!(txn.network_runtxp(&mut session, &mut MockResolver));
assert!(String::from_utf8_lossy(&session.output).contains("not an integer"));
}
}