package db import ( . "code.hoteas.com/golang/hotime/common" "code.hoteas.com/golang/hotime/log" "database/sql" "sync" "testing" ) func newTestDB(t *testing.T) *HoTimeDB { t.Helper() sqlDB, err := sql.Open("sqlite3", ":memory:") if err != nil { t.Fatal(err) } db := &HoTimeDB{ DB: sqlDB, Type: "sqlite3", Dialect: &SQLiteDialect{}, Log: log.NewTestLogger(), } _, execErr := sqlDB.Exec("CREATE TABLE test_item (id INTEGER PRIMARY KEY, name TEXT, value INTEGER)") if execErr != nil { t.Fatal(execErr) } _, execErr = sqlDB.Exec("INSERT INTO test_item (id, name, value) VALUES (1, 'foo', 100)") if execErr != nil { t.Fatal(execErr) } return db } func newTestDBWithTestTx(t *testing.T) *HoTimeDB { t.Helper() db := newTestDB(t) tx, err := db.DB.Begin() if err != nil { t.Fatal(err) } db.testTx = tx db.testMu = &sync.Mutex{} return db } func TestAction_TxFailedForceRollback(t *testing.T) { db := newTestDB(t) defer db.DB.Close() result := db.Action(func(tx HoTimeDB) (isSuccess bool) { tx.Update("test_item", Map{"value": 999}, Map{"id": 1}) tx.Exec("THIS IS INVALID SQL THAT WILL FAIL") return true }) if result != false { t.Fatalf("Action should return false when SQL error occurred, got true") } row := db.Get("test_item", "value", Map{"id": 1}) if row == nil { t.Fatal("failed to read test_item") } val := row.GetCeilInt64("value") if val != 100 { t.Fatalf("value should be rolled back to 100, got %d", val) } } func TestAction_NormalCommit(t *testing.T) { db := newTestDB(t) defer db.DB.Close() result := db.Action(func(tx HoTimeDB) (isSuccess bool) { tx.Update("test_item", Map{"value": 200}, Map{"id": 1}) return true }) if result != true { t.Fatalf("Action should return true on success, got false") } row := db.Get("test_item", "value", Map{"id": 1}) if row == nil { t.Fatal("failed to read test_item") } val := row.GetCeilInt64("value") if val != 200 { t.Fatalf("value should be committed to 200, got %d", val) } } func TestAction_NormalRollback(t *testing.T) { db := newTestDB(t) defer db.DB.Close() result := db.Action(func(tx HoTimeDB) (isSuccess bool) { tx.Update("test_item", Map{"value": 300}, Map{"id": 1}) return false }) if result != false { t.Fatalf("Action should return false, got true") } row := db.Get("test_item", "value", Map{"id": 1}) if row == nil { t.Fatal("failed to read test_item") } val := row.GetCeilInt64("value") if val != 100 { t.Fatalf("value should be rolled back to 100, got %d", val) } } func TestAction_SqlErrorThenReturnTrue_MustRollback(t *testing.T) { db := newTestDB(t) defer db.DB.Close() result := db.Action(func(tx HoTimeDB) (isSuccess bool) { tx.Update("test_item", Map{"value": 500}, Map{"id": 1}) tx.Exec("INSERT INTO nonexistent_table (x) VALUES (1)") tx.Update("test_item", Map{"value": 600}, Map{"id": 1}) return true }) if result != false { t.Fatalf("Action should return false when SQL error occurred mid-transaction, got true") } row := db.Get("test_item", "value", Map{"id": 1}) if row == nil { t.Fatal("failed to read test_item") } val := row.GetCeilInt64("value") if val != 100 { t.Fatalf("value should be rolled back to 100 (original), got %d", val) } } // --- testTx (SAVEPOINT) 模式测试 --- func TestAction_TestTx_SqlErrorForceRollbackSavepoint(t *testing.T) { db := newTestDBWithTestTx(t) defer db.testTx.Rollback() defer db.DB.Close() db.testTx.Exec("UPDATE test_item SET value = 100 WHERE id = 1") result := db.Action(func(tx HoTimeDB) (isSuccess bool) { tx.Update("test_item", Map{"value": 888}, Map{"id": 1}) tx.Exec("THIS IS INVALID SQL") return true }) if result != false { t.Fatalf("Action (testTx mode) should return false when SQL error occurred, got true") } db.testMu.Lock() rows, err := db.testTx.Query("SELECT value FROM test_item WHERE id = 1") db.testMu.Unlock() if err != nil { t.Fatal(err) } defer rows.Close() if !rows.Next() { t.Fatal("no row found") } var val int64 rows.Scan(&val) if val != 100 { t.Fatalf("value should be rolled back to 100 via SAVEPOINT, got %d", val) } } func TestAction_TestTx_NormalCommitSavepoint(t *testing.T) { db := newTestDBWithTestTx(t) defer db.testTx.Rollback() defer db.DB.Close() result := db.Action(func(tx HoTimeDB) (isSuccess bool) { tx.Update("test_item", Map{"value": 777}, Map{"id": 1}) return true }) if result != true { t.Fatalf("Action (testTx mode) should return true on success, got false") } db.testMu.Lock() rows, err := db.testTx.Query("SELECT value FROM test_item WHERE id = 1") db.testMu.Unlock() if err != nil { t.Fatal(err) } defer rows.Close() if !rows.Next() { t.Fatal("no row found") } var val int64 rows.Scan(&val) if val != 777 { t.Fatalf("value should be 777 after RELEASE SAVEPOINT, got %d", val) } } func TestAction_TestTx_NormalRollbackSavepoint(t *testing.T) { db := newTestDBWithTestTx(t) defer db.testTx.Rollback() defer db.DB.Close() db.testTx.Exec("UPDATE test_item SET value = 100 WHERE id = 1") result := db.Action(func(tx HoTimeDB) (isSuccess bool) { tx.Update("test_item", Map{"value": 666}, Map{"id": 1}) return false }) if result != false { t.Fatalf("Action (testTx mode) should return false, got true") } db.testMu.Lock() rows, err := db.testTx.Query("SELECT value FROM test_item WHERE id = 1") db.testMu.Unlock() if err != nil { t.Fatal(err) } defer rows.Close() if !rows.Next() { t.Fatal("no row found") } var val int64 rows.Scan(&val) if val != 100 { t.Fatalf("value should be rolled back to 100 via SAVEPOINT, got %d", val) } }