use super::{AllHeaderTy, Encode, ALL_HEADERS_LEN_TX};
use bytes::{BufMut, BytesMut};
use std::borrow::Cow;
uint_enum! {
#[repr(u16)]
pub enum TransactionManagerRequestType {
GetDtcAddress = 0,
Propagate = 1,
Begin = 5,
Promote = 6,
Commit = 7,
Rollback = 8,
Save = 9,
}
}
uint_enum! {
#[repr(u8)]
pub enum IsolationLevel {
Unspecified = 0x00,
ReadUncommitted = 0x01,
ReadCommitted = 0x02,
RepeatableRead = 0x03,
Serializable = 0x04,
Snapshot = 0x05,
}
}
#[derive(Debug, Clone)]
pub struct TransactionManagerRequest<'a> {
transaction_desc: [u8; 8],
body: TransactionRequestBody<'a>,
}
#[derive(Debug, Clone)]
enum TransactionRequestBody<'a> {
Begin {
isolation_level: IsolationLevel,
name: Cow<'a, str>,
},
Commit {
name: Cow<'a, str>,
},
Rollback {
name: Cow<'a, str>,
},
Save {
name: Cow<'a, str>,
},
}
impl<'a> TransactionManagerRequest<'a> {
pub fn begin(
transaction_desc: [u8; 8],
isolation_level: IsolationLevel,
name: impl Into<Cow<'a, str>>,
) -> Self {
Self {
transaction_desc,
body: TransactionRequestBody::Begin {
isolation_level,
name: name.into(),
},
}
}
pub fn commit(transaction_desc: [u8; 8], name: impl Into<Cow<'a, str>>) -> Self {
Self {
transaction_desc,
body: TransactionRequestBody::Commit { name: name.into() },
}
}
pub fn rollback(transaction_desc: [u8; 8], name: impl Into<Cow<'a, str>>) -> Self {
Self {
transaction_desc,
body: TransactionRequestBody::Rollback { name: name.into() },
}
}
pub fn save(transaction_desc: [u8; 8], name: impl Into<Cow<'a, str>>) -> Self {
Self {
transaction_desc,
body: TransactionRequestBody::Save { name: name.into() },
}
}
fn request_type(&self) -> TransactionManagerRequestType {
match self.body {
TransactionRequestBody::Begin { .. } => TransactionManagerRequestType::Begin,
TransactionRequestBody::Commit { .. } => TransactionManagerRequestType::Commit,
TransactionRequestBody::Rollback { .. } => TransactionManagerRequestType::Rollback,
TransactionRequestBody::Save { .. } => TransactionManagerRequestType::Save,
}
}
}
fn encode_b_varchar(dst: &mut BytesMut, s: &str) {
let units: Vec<u16> = s.encode_utf16().collect();
dst.put_u8(units.len() as u8);
for unit in units {
dst.put_u16_le(unit);
}
}
impl<'a> Encode<BytesMut> for TransactionManagerRequest<'a> {
fn encode(self, dst: &mut BytesMut) -> crate::Result<()> {
dst.put_u32_le(ALL_HEADERS_LEN_TX as u32);
dst.put_u32_le(ALL_HEADERS_LEN_TX as u32 - 4);
dst.put_u16_le(AllHeaderTy::TransactionDescriptor as u16);
dst.put_slice(&self.transaction_desc);
dst.put_u32_le(1);
dst.put_u16_le(self.request_type() as u16);
match self.body {
TransactionRequestBody::Begin {
isolation_level,
name,
} => {
dst.put_u8(isolation_level as u8);
encode_b_varchar(dst, &name);
}
TransactionRequestBody::Commit { name } | TransactionRequestBody::Rollback { name } => {
encode_b_varchar(dst, &name);
dst.put_u8(0);
}
TransactionRequestBody::Save { name } => {
encode_b_varchar(dst, &name);
}
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
fn all_headers() -> Vec<u8> {
let mut v = Vec::new();
v.extend_from_slice(&(ALL_HEADERS_LEN_TX as u32).to_le_bytes());
v.extend_from_slice(&(ALL_HEADERS_LEN_TX as u32 - 4).to_le_bytes());
v.extend_from_slice(&(AllHeaderTy::TransactionDescriptor as u16).to_le_bytes());
v.extend_from_slice(&[1, 2, 3, 4, 5, 6, 7, 8]);
v.extend_from_slice(&1u32.to_le_bytes());
v
}
#[test]
fn encodes_begin_request() {
let desc = [1, 2, 3, 4, 5, 6, 7, 8];
let req = TransactionManagerRequest::begin(desc, IsolationLevel::ReadCommitted, "");
let mut buf = BytesMut::new();
req.encode(&mut buf).unwrap();
let mut expected = all_headers();
expected.extend_from_slice(&(TransactionManagerRequestType::Begin as u16).to_le_bytes());
expected.push(IsolationLevel::ReadCommitted as u8); expected.push(0);
assert_eq!(&buf[..], &expected[..]);
}
#[test]
fn encodes_begin_request_with_name() {
let desc = [1, 2, 3, 4, 5, 6, 7, 8];
let req = TransactionManagerRequest::begin(desc, IsolationLevel::Serializable, "tx");
let mut buf = BytesMut::new();
req.encode(&mut buf).unwrap();
let mut expected = all_headers();
expected.extend_from_slice(&(TransactionManagerRequestType::Begin as u16).to_le_bytes());
expected.push(IsolationLevel::Serializable as u8);
expected.push(2); expected.extend_from_slice(&b't'.to_le_bytes());
expected.push(0);
expected.extend_from_slice(&b'x'.to_le_bytes());
expected.push(0);
assert_eq!(&buf[..], &expected[..]);
}
#[test]
fn encodes_commit_request() {
let desc = [1, 2, 3, 4, 5, 6, 7, 8];
let req = TransactionManagerRequest::commit(desc, "");
let mut buf = BytesMut::new();
req.encode(&mut buf).unwrap();
let mut expected = all_headers();
expected.extend_from_slice(&(TransactionManagerRequestType::Commit as u16).to_le_bytes());
expected.push(0); expected.push(0);
assert_eq!(&buf[..], &expected[..]);
}
#[test]
fn encodes_rollback_request() {
let desc = [8, 7, 6, 5, 4, 3, 2, 1];
let req = TransactionManagerRequest::rollback(desc, "");
let mut buf = BytesMut::new();
req.encode(&mut buf).unwrap();
let mut expected = Vec::new();
expected.extend_from_slice(&(ALL_HEADERS_LEN_TX as u32).to_le_bytes());
expected.extend_from_slice(&(ALL_HEADERS_LEN_TX as u32 - 4).to_le_bytes());
expected.extend_from_slice(&(AllHeaderTy::TransactionDescriptor as u16).to_le_bytes());
expected.extend_from_slice(&desc);
expected.extend_from_slice(&1u32.to_le_bytes());
expected.extend_from_slice(&(TransactionManagerRequestType::Rollback as u16).to_le_bytes());
expected.push(0); expected.push(0);
assert_eq!(&buf[..], &expected[..]);
}
#[test]
fn encodes_save_request() {
let desc = [1, 2, 3, 4, 5, 6, 7, 8];
let req = TransactionManagerRequest::save(desc, "sp1");
let mut buf = BytesMut::new();
req.encode(&mut buf).unwrap();
let mut expected = all_headers();
expected.extend_from_slice(&(TransactionManagerRequestType::Save as u16).to_le_bytes());
expected.push(3); for unit in "sp1".encode_utf16() {
expected.extend_from_slice(&unit.to_le_bytes());
}
assert_eq!(&buf[..], &expected[..]);
}
}