#![cfg(test)]
#[path = "common/mod.rs"]
mod common;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::Mutex;
use sz_orm_core::{Connection, TransactOptions, Transaction, TransactionState, TxError};
use common::{InMemoryDb, TransactionalConnection};
fn make_conn() -> (TransactionalConnection, Arc<Mutex<InMemoryDb>>) {
let db = Arc::new(Mutex::new(InMemoryDb::new()));
let conn = TransactionalConnection::new(db.clone());
(conn, db)
}
async fn seed_users(db: &Arc<Mutex<InMemoryDb>>) {
let mut d = db.lock().await;
d.create_table("users");
let mut row = std::collections::HashMap::new();
row.insert("id".to_string(), sz_orm_core::Value::I64(1));
row.insert(
"name".to_string(),
sz_orm_core::Value::String("Alice".to_string()),
);
row.insert("age".to_string(), sz_orm_core::Value::I64(30));
d.insert("users", row);
}
#[tokio::test]
async fn test_l3_t1_basic_commit_persists_data() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
conn.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 2, "事务内应可见新插入的行");
}
conn.commit().await.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 2, "commit 后数据应持久化");
let bob = d.find_where("users", "id", &sz_orm_core::Value::I64(2));
assert!(bob.is_some(), "Bob 的行应存在");
assert_eq!(
bob.unwrap().get("name"),
Some(&sz_orm_core::Value::String("Bob".to_string()))
);
}
}
#[tokio::test]
async fn test_l3_t2_basic_rollback_discards_data() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
conn.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 2, "事务内应可见新插入的行");
}
conn.rollback().await.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 1, "rollback 后新插入的行应被丢弃");
assert!(
d.find_where("users", "id", &sz_orm_core::Value::I64(2))
.is_none(),
"Bob 的行不应存在"
);
}
}
#[tokio::test]
async fn test_l3_t3_multi_statement_commit() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
conn.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
conn.execute("UPDATE users SET age = 31 WHERE id = 1")
.await
.unwrap();
conn.execute("DELETE FROM users WHERE id = 999")
.await
.unwrap();
conn.commit().await.unwrap();
let d = db.lock().await;
assert_eq!(d.count("users"), 2, "应有 2 行(Alice + Bob)");
let alice = d.find_where("users", "id", &sz_orm_core::Value::I64(1));
assert_eq!(
alice.unwrap().get("age"),
Some(&sz_orm_core::Value::I64(31)),
"Alice 的 age 应被更新为 31"
);
}
#[tokio::test]
async fn test_l3_t4_multi_statement_rollback() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
let original_snapshot = db.lock().await.snapshot("users");
conn.begin_transaction().await.unwrap();
conn.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
conn.execute("UPDATE users SET age = 99 WHERE id = 1")
.await
.unwrap();
conn.execute("DELETE FROM users WHERE id = 1")
.await
.unwrap();
conn.rollback().await.unwrap();
let d = db.lock().await;
let current = d.snapshot("users");
assert_eq!(
current.len(),
original_snapshot.len(),
"rollback 后行数应恢复原始值"
);
let alice = d.find_where("users", "id", &sz_orm_core::Value::I64(1));
assert_eq!(
alice.unwrap().get("age"),
Some(&sz_orm_core::Value::I64(30)),
"Alice 的 age 应恢复为原始值 30"
);
}
#[tokio::test]
async fn test_l3_t5_savepoint_rollback() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
conn.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
let mut tx = take_conn_into_tx(conn);
let sp = tx.savepoint().await.unwrap();
tx.execute("INSERT INTO users (id, name, age) VALUES (3, 'Charlie', 40)")
.await
.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 3, "保存点后应有 3 行");
}
tx.rollback_to_savepoint(&sp).await.unwrap();
tx.release_savepoint(&sp).await.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 2, "回滚到保存点后应有 2 行");
assert!(
d.find_where("users", "id", &sz_orm_core::Value::I64(3))
.is_none(),
"Charlie 的行应被回滚"
);
}
tx.commit().await.unwrap();
let d = db.lock().await;
assert_eq!(d.count("users"), 2, "commit 后应有 2 行(Alice + Bob)");
assert!(
d.find_where("users", "id", &sz_orm_core::Value::I64(2))
.is_some(),
"Bob 的行应保留"
);
}
#[tokio::test]
async fn test_l3_t6_nested_savepoints() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
let mut tx = take_conn_into_tx(conn);
tx.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
let sp1 = tx.savepoint().await.unwrap();
tx.execute("INSERT INTO users (id, name, age) VALUES (3, 'Charlie', 40)")
.await
.unwrap();
let sp2 = tx.savepoint().await.unwrap();
tx.execute("INSERT INTO users (id, name, age) VALUES (4, 'Dave', 50)")
.await
.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 4, "sp2 后应有 4 行");
}
tx.rollback_to_savepoint(&sp2).await.unwrap();
tx.release_savepoint(&sp2).await.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 3, "回滚到 sp2 后应有 3 行");
assert!(
d.find_where("users", "id", &sz_orm_core::Value::I64(4))
.is_none(),
"Dave 应被回滚"
);
}
tx.rollback_to_savepoint(&sp1).await.unwrap();
tx.release_savepoint(&sp1).await.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 2, "回滚到 sp1 后应有 2 行");
assert!(
d.find_where("users", "id", &sz_orm_core::Value::I64(3))
.is_none(),
"Charlie 应被回滚"
);
}
tx.commit().await.unwrap();
let d = db.lock().await;
assert_eq!(d.count("users"), 2, "commit 后应有 2 行(Alice + Bob)");
}
#[tokio::test]
async fn test_l3_t7_release_savepoint_then_rollback_all() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
let mut tx = take_conn_into_tx(conn);
let sp = tx.savepoint().await.unwrap();
tx.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
tx.release_savepoint(&sp).await.unwrap();
tx.rollback().await.unwrap();
let d = db.lock().await;
assert_eq!(d.count("users"), 1, "rollback 后应只有原始的 Alice");
assert!(
d.find_where("users", "id", &sz_orm_core::Value::I64(2))
.is_none(),
"Bob 应被回滚"
);
}
#[tokio::test]
async fn test_l3_t8_transaction_timeout_rollback() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
let opts = TransactOptions::default().with_timeout(Duration::from_millis(50));
let mut tx = Transaction::new(Box::new(conn), opts);
tx.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(100)).await;
let result = tx.commit().await;
assert!(result.is_err(), "超时后 commit 应失败");
match result.unwrap_err() {
TxError::CommitFailed(msg) => {
assert!(msg.contains("timeout"), "错误信息应包含 timeout");
}
other => panic!("期望 CommitFailed(timeout),实际: {:?}", other),
}
assert_eq!(tx.state(), TransactionState::RolledBack);
let d = db.lock().await;
assert_eq!(d.count("users"), 1, "超时回滚后应只有原始的 Alice");
}
#[tokio::test]
async fn test_l3_t9_max_nesting_depth_enforced() {
let (mut conn, _db) = make_conn();
conn.begin_transaction().await.unwrap();
let opts = TransactOptions::default().with_max_nesting_depth(3);
let mut tx = Transaction::new(Box::new(conn), opts);
for i in 1..=3 {
let sp = tx.savepoint().await.unwrap();
assert_eq!(sp, format!("sp_{}", i));
}
let result = tx.savepoint().await;
assert!(result.is_err());
match result.unwrap_err() {
TxError::MaxNestingDepthExceeded {
current_depth,
max_depth,
} => {
assert_eq!(current_depth, 4);
assert_eq!(max_depth, 3);
}
other => panic!("期望 MaxNestingDepthExceeded,实际: {:?}", other),
}
}
#[tokio::test]
async fn test_l3_t10_drop_auto_rollback() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
let mut tx = Transaction::new(Box::new(conn), TransactOptions::default());
tx.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
drop(tx);
tokio::time::sleep(Duration::from_millis(100)).await;
let d = db.lock().await;
assert_eq!(d.count("users"), 1, "Drop 自动回滚后应只有原始的 Alice");
assert!(
d.find_where("users", "id", &sz_orm_core::Value::I64(2))
.is_none(),
"Bob 应被 Drop 自动回滚"
);
}
#[tokio::test]
async fn test_l3_t11_commit_failure_rolls_back() {
use std::future::Future;
use std::pin::Pin;
use sz_orm_core::DbError;
struct FailingCommitConnection {
inner: TransactionalConnection,
}
impl Connection for FailingCommitConnection {
fn execute<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
self.inner.execute(sql)
}
fn query<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<
Box<
dyn Future<
Output = Result<
Vec<std::collections::HashMap<String, sz_orm_core::Value>>,
DbError,
>,
> + Send
+ 'a,
>,
> {
self.inner.query(sql)
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
self.inner.begin_transaction()
}
fn commit<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
Err(DbError::Internal("injected commit failure".to_string()))
})
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
self.inner.rollback()
}
fn is_connected(&self) -> bool {
self.inner.is_connected()
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
self.inner.ping()
}
fn close<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
self.inner.close()
}
}
let db = Arc::new(Mutex::new(InMemoryDb::new()));
seed_users(&db).await;
let inner = TransactionalConnection::new(db.clone());
let mut conn = FailingCommitConnection { inner };
conn.begin_transaction().await.unwrap();
conn.execute("INSERT INTO users (id, name, age) VALUES (2, 'Bob', 25)")
.await
.unwrap();
let mut tx = Transaction::new(Box::new(conn), TransactOptions::default());
let result = tx.commit().await;
assert!(result.is_err(), "commit 应失败");
assert_eq!(tx.state(), TransactionState::Active);
tx.rollback().await.unwrap();
let d = db.lock().await;
assert_eq!(
d.count("users"),
1,
"commit 失败 + rollback 后应只有原始 Alice"
);
}
#[tokio::test]
async fn test_l3_t12_isolation_level_options() {
use sz_orm_core::IsolationLevel;
let variants = vec![
IsolationLevel::ReadUncommitted,
IsolationLevel::ReadCommitted,
IsolationLevel::RepeatableRead,
IsolationLevel::Serializable,
IsolationLevel::Snapshot,
];
for level in variants {
let (conn, _db) = make_conn();
let opts = TransactOptions::default()
.with_isolation(level.clone())
.read_only()
.with_timeout(Duration::from_secs(30));
let tx = Transaction::new(Box::new(conn), opts);
assert_eq!(
tx.options().isolation_level,
Some(level),
"隔离级别 round-trip 应保持值不变"
);
assert!(tx.options().read_only, "read_only 应为 true");
assert_eq!(
tx.options().timeout,
Some(Duration::from_secs(30)),
"timeout 应为 30s"
);
}
let (conn, _db) = make_conn();
let tx = Transaction::new(Box::new(conn), TransactOptions::default());
assert_eq!(
tx.options().isolation_level,
None,
"default opts 不应设置隔离级别"
);
}
#[tokio::test]
async fn test_l3_t13_propagation_behavior_options() {
use sz_orm_core::PropagationBehavior;
let variants = vec![
PropagationBehavior::Required,
PropagationBehavior::Mandatory,
PropagationBehavior::Never,
PropagationBehavior::Supports,
PropagationBehavior::RequiresNew,
PropagationBehavior::Nested,
];
for behavior in variants {
let (conn, _db) = make_conn();
let opts = TransactOptions::default().with_propagation(behavior);
let tx = Transaction::new(Box::new(conn), opts);
assert_eq!(
tx.options().propagation,
behavior,
"传播行为 round-trip 应保持值不变"
);
}
let (conn, _db) = make_conn();
let tx = Transaction::new(Box::new(conn), TransactOptions::default());
assert_eq!(
tx.options().propagation,
PropagationBehavior::Required,
"default 传播行为应为 Required"
);
}
#[tokio::test]
async fn test_l3_t14_multi_table_transaction_rollback() {
let (mut conn, db) = make_conn();
{
let mut d = db.lock().await;
d.create_table("users");
d.create_table("orders");
let mut user = std::collections::HashMap::new();
user.insert("id".to_string(), sz_orm_core::Value::I64(1));
user.insert(
"name".to_string(),
sz_orm_core::Value::String("Alice".to_string()),
);
d.insert("users", user);
let mut order = std::collections::HashMap::new();
order.insert("id".to_string(), sz_orm_core::Value::I64(100));
order.insert("user_id".to_string(), sz_orm_core::Value::I64(1));
order.insert("amount".to_string(), sz_orm_core::Value::I64(99));
d.insert("orders", order);
}
conn.begin_transaction().await.unwrap();
let mut tx = take_conn_into_tx(conn);
tx.execute("INSERT INTO users (id, name) VALUES (2, 'Bob')")
.await
.unwrap();
tx.execute("INSERT INTO orders (id, user_id, amount) VALUES (101, 2, 50)")
.await
.unwrap();
tx.execute("UPDATE orders SET amount = 199 WHERE id = 100")
.await
.unwrap();
{
let d = db.lock().await;
assert_eq!(d.count("users"), 2);
assert_eq!(d.count("orders"), 2);
}
tx.rollback().await.unwrap();
let d = db.lock().await;
assert_eq!(d.count("users"), 1, "users 表应恢复原始 1 行");
assert_eq!(d.count("orders"), 1, "orders 表应恢复原始 1 行");
let order = d.find_where("orders", "id", &sz_orm_core::Value::I64(100));
assert_eq!(
order.unwrap().get("amount"),
Some(&sz_orm_core::Value::I64(99)),
"orders 表的 amount 应恢复为原始值 99"
);
}
#[tokio::test]
async fn test_l3_t15_update_rollback_restores_original() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
let mut tx = take_conn_into_tx(conn);
tx.execute("UPDATE users SET age = 40 WHERE id = 1")
.await
.unwrap();
tx.execute("UPDATE users SET name = 'Alice2' WHERE id = 1")
.await
.unwrap();
{
let d = db.lock().await;
let alice = d.find_where("users", "id", &sz_orm_core::Value::I64(1));
let alice = alice.unwrap();
assert_eq!(alice.get("age"), Some(&sz_orm_core::Value::I64(40)));
assert_eq!(
alice.get("name"),
Some(&sz_orm_core::Value::String("Alice2".to_string()))
);
}
tx.rollback().await.unwrap();
let d = db.lock().await;
let alice = d.find_where("users", "id", &sz_orm_core::Value::I64(1));
let alice = alice.unwrap();
assert_eq!(
alice.get("age"),
Some(&sz_orm_core::Value::I64(30)),
"age 应恢复为 30"
);
assert_eq!(
alice.get("name"),
Some(&sz_orm_core::Value::String("Alice".to_string())),
"name 应恢复为 Alice"
);
}
#[tokio::test]
async fn test_l3_t16_delete_rollback_restores_data() {
let (mut conn, db) = make_conn();
seed_users(&db).await;
conn.begin_transaction().await.unwrap();
let mut tx = take_conn_into_tx(conn);
let affected = tx.execute("DELETE FROM users WHERE id = 1").await.unwrap();
assert_eq!(affected, 1, "应删除 1 行");
{
let d = db.lock().await;
assert_eq!(d.count("users"), 0, "事务内 users 应为空");
}
tx.rollback().await.unwrap();
let d = db.lock().await;
assert_eq!(d.count("users"), 1, "rollback 后 users 应恢复 1 行");
let alice = d.find_where("users", "id", &sz_orm_core::Value::I64(1));
assert!(alice.is_some(), "Alice 应恢复");
}
fn take_conn_into_tx(conn: TransactionalConnection) -> Transaction {
Transaction::new(Box::new(conn), TransactOptions::default())
}