diff --git a/dtmcli/barrier.go b/dtmcli/barrier.go index 1fcf61f..2235471 100644 --- a/dtmcli/barrier.go +++ b/dtmcli/barrier.go @@ -20,11 +20,13 @@ type BarrierBusiFunc func(tx *sql.Tx) error // BranchBarrier every branch info type BranchBarrier struct { - TransType string - Gid string - BranchID string - Op string - BarrierID int + TransType string + Gid string + BranchID string + Op string + BarrierID int + DBType string // DBTypeMysql | DBTypePostgres + BarrierTableName string } func (bb *BranchBarrier) String() string { @@ -70,8 +72,8 @@ func (bb *BranchBarrier) Call(tx *sql.Tx, busiCall BarrierBusiFunc) (rerr error) dtmimp.OpCompensate: dtmimp.OpAction, }[bb.Op] - originAffected, oerr := dtmimp.InsertBarrier(tx, bb.TransType, bb.Gid, bb.BranchID, originOp, bid, bb.Op) - currentAffected, rerr := dtmimp.InsertBarrier(tx, bb.TransType, bb.Gid, bb.BranchID, bb.Op, bid, bb.Op) + originAffected, oerr := dtmimp.InsertBarrier(tx, bb.TransType, bb.Gid, bb.BranchID, originOp, bid, bb.Op, bb.DBType, bb.BarrierTableName) + currentAffected, rerr := dtmimp.InsertBarrier(tx, bb.TransType, bb.Gid, bb.BranchID, bb.Op, bid, bb.Op, bb.DBType, bb.BarrierTableName) logger.Debugf("originAffected: %d currentAffected: %d", originAffected, currentAffected) if rerr == nil && bb.Op == dtmimp.MsgDoOp && currentAffected == 0 { // for msg's DoAndSubmit, repeated insert should be rejected. @@ -103,7 +105,7 @@ func (bb *BranchBarrier) CallWithDB(db *sql.DB, busiCall BarrierBusiFunc) error // QueryPrepared queries prepared data func (bb *BranchBarrier) QueryPrepared(db *sql.DB) error { - _, err := dtmimp.InsertBarrier(db, bb.TransType, bb.Gid, dtmimp.MsgDoBranch0, dtmimp.MsgDoOp, dtmimp.MsgDoBarrier1, dtmimp.OpRollback) + _, err := dtmimp.InsertBarrier(db, bb.TransType, bb.Gid, dtmimp.MsgDoBranch0, dtmimp.MsgDoOp, dtmimp.MsgDoBarrier1, dtmimp.OpRollback, bb.DBType, bb.BarrierTableName) var reason string if err == nil { sql := fmt.Sprintf("select reason from %s where gid=? and branch_id=? and op=? and barrier_id=?", dtmimp.BarrierTableName) diff --git a/dtmcli/dtmimp/db_special.go b/dtmcli/dtmimp/db_special.go index b4db7fc..fed04f6 100644 --- a/dtmcli/dtmimp/db_special.go +++ b/dtmcli/dtmimp/db_special.go @@ -75,8 +75,8 @@ func init() { } // GetDBSpecial get DBSpecial for currentDBType -func GetDBSpecial() DBSpecial { - return dbSpecials[currentDBType] +func GetDBSpecial(dbType string) DBSpecial { + return dbSpecials[dbType] } // SetCurrentDBType set currentDBType diff --git a/dtmcli/dtmimp/db_special_test.go b/dtmcli/dtmimp/db_special_test.go index 3966cd2..7003109 100644 --- a/dtmcli/dtmimp/db_special_test.go +++ b/dtmcli/dtmimp/db_special_test.go @@ -18,13 +18,13 @@ func TestDBSpecial(t *testing.T) { SetCurrentDBType("no-driver") })) SetCurrentDBType(DBTypeMysql) - sp := GetDBSpecial() + sp := GetDBSpecial(DBTypeMysql) assert.Equal(t, "? ?", sp.GetPlaceHoldSQL("? ?")) assert.Equal(t, "xa start 'xa1'", sp.GetXaSQL("start", "xa1")) assert.Equal(t, "insert ignore into a(f) values(?)", sp.GetInsertIgnoreTemplate("a(f) values(?)", "c")) SetCurrentDBType(DBTypePostgres) - sp = GetDBSpecial() + sp = GetDBSpecial(DBTypeMysql) assert.Equal(t, "$1 $2", sp.GetPlaceHoldSQL("? ?")) assert.Equal(t, "begin", sp.GetXaSQL("start", "xa1")) assert.Equal(t, "insert into a(f) values(?) on conflict ON CONSTRAINT c do nothing", sp.GetInsertIgnoreTemplate("a(f) values(?)", "c")) diff --git a/dtmcli/dtmimp/trans_xa_base.go b/dtmcli/dtmimp/trans_xa_base.go index 50bb9c8..bd15c87 100644 --- a/dtmcli/dtmimp/trans_xa_base.go +++ b/dtmcli/dtmimp/trans_xa_base.go @@ -18,14 +18,14 @@ func XaHandlePhase2(gid string, dbConf DBConf, branchID string, op string) error return err } xaID := gid + "-" + branchID - _, err = DBExec(db, GetDBSpecial().GetXaSQL(op, xaID)) + _, err = DBExec(dbConf.Driver, db, GetDBSpecial(dbConf.Driver).GetXaSQL(op, xaID)) if err != nil && (strings.Contains(err.Error(), "XAER_NOTA") || strings.Contains(err.Error(), "does not exist")) { // Repeat commit/rollback with the same id, report this error, ignore err = nil } if op == OpRollback && err == nil { // rollback insert a row after prepare. no-error means prepare has finished. - _, err = InsertBarrier(db, "xa", gid, branchID, OpAction, XaBarrier1, op) + _, err = InsertBarrier(db, "xa", gid, branchID, OpAction, XaBarrier1, op, dbConf.Driver, "TODO") } return err } @@ -39,20 +39,20 @@ func XaHandleLocalTrans(xa *TransBase, dbConf DBConf, cb func(*sql.DB) error) (r } defer func() { _ = db.Close() }() defer DeferDo(&rerr, func() error { - _, err := DBExec(db, GetDBSpecial().GetXaSQL("prepare", xaBranch)) + _, err := DBExec(dbConf.Driver, db, GetDBSpecial(dbConf.Driver).GetXaSQL("prepare", xaBranch)) return err }, func() error { return nil }) - _, rerr = DBExec(db, GetDBSpecial().GetXaSQL("start", xaBranch)) + _, rerr = DBExec(dbConf.Driver, db, GetDBSpecial(dbConf.Driver).GetXaSQL("start", xaBranch)) if rerr != nil { return } defer func() { - _, _ = DBExec(db, GetDBSpecial().GetXaSQL("end", xaBranch)) + _, _ = DBExec(dbConf.Driver, db, GetDBSpecial(dbConf.Driver).GetXaSQL("end", xaBranch)) }() // prepare and rollback both insert a row - _, rerr = InsertBarrier(db, xa.TransType, xa.Gid, xa.BranchID, OpAction, XaBarrier1, OpAction) + _, rerr = InsertBarrier(db, xa.TransType, xa.Gid, xa.BranchID, OpAction, XaBarrier1, OpAction, dbConf.Driver, "TODO") if rerr == nil { rerr = cb(db) } diff --git a/dtmcli/dtmimp/utils.go b/dtmcli/dtmimp/utils.go index 7220b30..1e5607a 100644 --- a/dtmcli/dtmimp/utils.go +++ b/dtmcli/dtmimp/utils.go @@ -187,12 +187,12 @@ func XaDB(conf DBConf) (*sql.DB, error) { } // DBExec use raw db to exec -func DBExec(db DB, sql string, values ...interface{}) (affected int64, rerr error) { +func DBExec(dbType string, db DB, sql string, values ...interface{}) (affected int64, rerr error) { if sql == "" { return 0, nil } began := time.Now() - sql = GetDBSpecial().GetPlaceHoldSQL(sql) + sql = GetDBSpecial(dbType).GetPlaceHoldSQL(sql) r, rerr := db.Exec(sql, values...) used := time.Since(began) / time.Millisecond if rerr == nil { @@ -262,10 +262,16 @@ func EscapeGet(qs url.Values, key string) string { } // InsertBarrier insert a record to barrier -func InsertBarrier(tx DB, transType string, gid string, branchID string, op string, barrierID string, reason string) (int64, error) { +func InsertBarrier(tx DB, transType string, gid string, branchID string, op string, barrierID string, reason string, dbType string, barrierTableName string) (int64, error) { if op == "" { return 0, nil } - sql := GetDBSpecial().GetInsertIgnoreTemplate(BarrierTableName+"(trans_type, gid, branch_id, op, barrier_id, reason) values(?,?,?,?,?,?)", "uniq_barrier") - return DBExec(tx, sql, transType, gid, branchID, op, barrierID, reason) + if dbType == "" { + dbType = currentDBType + } + if barrierTableName == "" { + barrierTableName = BarrierTableName + } + sql := GetDBSpecial(dbType).GetInsertIgnoreTemplate(barrierTableName+"(trans_type, gid, branch_id, op, barrier_id, reason) values(?,?,?,?,?,?)", "uniq_barrier") + return DBExec(dbType, tx, sql, transType, gid, branchID, op, barrierID, reason) } diff --git a/dtmutil/utils.go b/dtmutil/utils.go index 06b39a7..599419c 100644 --- a/dtmutil/utils.go +++ b/dtmutil/utils.go @@ -168,7 +168,7 @@ func RunSQLScript(conf dtmcli.DBConf, script string, skipDrop bool) { if s == "" || (skipDrop && strings.Contains(s, "drop")) { continue } - _, err = dtmimp.DBExec(con, s) + _, err = dtmimp.DBExec(conf.Driver, con, s) logger.FatalIfError(err) logger.Infof("sql scripts finished: %s", s) } diff --git a/test/busi/busi.go b/test/busi/busi.go index e31360c..7d24985 100644 --- a/test/busi/busi.go +++ b/test/busi/busi.go @@ -66,7 +66,7 @@ func sagaGrpcAdjustBalance(db dtmcli.DB, uid int, amount int64, result string) e if result == dtmcli.ResultFailure { return status.New(codes.Aborted, dtmcli.ResultFailure).Err() } - _, err := dtmimp.DBExec(db, "update dtm_busi.user_account set balance = balance + ? where user_id = ?", amount, uid) + _, err := dtmimp.DBExec(BusiConf.Driver, db, "update dtm_busi.user_account set balance = balance + ? where user_id = ?", amount, uid) return err } @@ -75,7 +75,7 @@ func SagaAdjustBalance(db dtmcli.DB, uid int, amount int, result string) error { if strings.Contains(result, dtmcli.ResultFailure) { return dtmcli.ErrFailure } - _, err := dtmimp.DBExec(db, "update dtm_busi.user_account set balance = balance + ? where user_id = ?", amount, uid) + _, err := dtmimp.DBExec(BusiConf.Driver, db, "update dtm_busi.user_account set balance = balance + ? where user_id = ?", amount, uid) return err } @@ -102,11 +102,10 @@ func SagaMongoAdjustBalance(ctx context.Context, mc *mongo.Client, uid int, amou return fmt.Errorf("balance not enough %w", dtmcli.ErrFailure) } return nil - } func tccAdjustTrading(db dtmcli.DB, uid int, amount int) error { - affected, err := dtmimp.DBExec(db, `update dtm_busi.user_account + affected, err := dtmimp.DBExec(BusiConf.Driver, db, `update dtm_busi.user_account set trading_balance=trading_balance+? where user_id=? and trading_balance + ? + balance >= 0`, amount, uid, amount) if err == nil && affected == 0 { @@ -116,7 +115,7 @@ func tccAdjustTrading(db dtmcli.DB, uid int, amount int) error { } func tccAdjustBalance(db dtmcli.DB, uid int, amount int) error { - affected, err := dtmimp.DBExec(db, `update dtm_busi.user_account + affected, err := dtmimp.DBExec(BusiConf.Driver, db, `update dtm_busi.user_account set trading_balance=trading_balance-?, balance=balance+? where user_id=?`, amount, amount, uid) if err == nil && affected == 0 { diff --git a/test/common_test.go b/test/common_test.go index 2a9b917..a28a144 100644 --- a/test/common_test.go +++ b/test/common_test.go @@ -33,12 +33,12 @@ func testSql(t *testing.T) { func testDbAlone(t *testing.T) { db, err := dtmimp.StandaloneDB(conf.Store.GetDBConf()) assert.Nil(t, err) - _, err = dtmimp.DBExec(db, "select 1") + _, err = dtmimp.DBExec(conf.Store.Driver, db, "select 1") assert.Equal(t, nil, err) - _, err = dtmimp.DBExec(db, "") + _, err = dtmimp.DBExec(conf.Store.Driver, db, "") assert.Equal(t, nil, err) db.Close() - _, err = dtmimp.DBExec(db, "select 1") + _, err = dtmimp.DBExec(conf.Store.Driver, db, "select 1") assert.NotEqual(t, nil, err) } diff --git a/test/xa_test.go b/test/xa_test.go index 4980a0a..ff17286 100644 --- a/test/xa_test.go +++ b/test/xa_test.go @@ -43,10 +43,10 @@ func TestXaDuplicate(t *testing.T) { sdb, err := dtmimp.StandaloneDB(busi.BusiConf) assert.Nil(t, err) if dtmcli.GetCurrentDBType() == dtmcli.DBTypeMysql { - _, err = dtmimp.DBExec(sdb, "xa recover") + _, err = dtmimp.DBExec(conf.Store.Driver, sdb, "xa recover") assert.Nil(t, err) } - _, err = dtmimp.DBExec(sdb, dtmimp.GetDBSpecial().GetXaSQL("commit", gid+"-01")) // simulate repeated request + _, err = dtmimp.DBExec(conf.Store.Driver, sdb, dtmimp.GetDBSpecial(conf.Store.Driver).GetXaSQL("commit", gid+"-01")) // simulate repeated request assert.Nil(t, err) return xa.CallBranch(req, busi.Busi+"/TransInXa") })