Files
hotime/db/transaction_test.go
T
hoteas f787421f82 fix(db): 修复事务中SQL错误时的回滚机制
- 添加txFailed标志跟踪事务内的SQL错误状态
- 在事务模式下遇到SQL错误时立即标记事务失败并阻止重试
- 修改Action方法确保SQL出错时强制回滚事务
- 为事务失败场景添加完整的单元测试覆盖
- 防止在事务已回滚的情况下继续执行后续操作导致数据不一致
2026-04-11 21:53:28 +08:00

136 lines
3.1 KiB
Go

package db
import (
. "code.hoteas.com/golang/hotime/common"
"database/sql"
"testing"
"github.com/sirupsen/logrus"
)
func newTestDB(t *testing.T) *HoTimeDB {
t.Helper()
sqlDB, err := sql.Open("sqlite3", ":memory:")
if err != nil {
t.Fatal(err)
}
logger := logrus.New()
logger.SetLevel(logrus.WarnLevel)
db := &HoTimeDB{
DB: sqlDB,
Type: "sqlite3",
Dialect: &SQLiteDialect{},
LastErr: &Error{Logger: logger},
Log: logger,
}
_, 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 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)
}
}