From eb9038f977bab1d981857a4169412b21e24f4017 Mon Sep 17 00:00:00 2001 From: yedf2 <120050102@qq.com> Date: Tue, 5 Oct 2021 17:00:06 +0800 Subject: [PATCH] partial: submit refactored --- .gitignore | 1 + app/main.go | 3 ++ bench/http.go | 124 +++++++++++++++++++++++++++++++++++++++++++++ bench/run.sh | 15 ++++++ dtmsvr/api.go | 25 ++++++--- dtmsvr/api_http.go | 5 +- dtmsvr/trans.go | 30 +++++------ test/saga_test.go | 3 +- 8 files changed, 179 insertions(+), 27 deletions(-) create mode 100644 bench/http.go create mode 100644 bench/run.sh diff --git a/.gitignore b/.gitignore index 1ab47fb..9304d6c 100644 --- a/.gitignore +++ b/.gitignore @@ -1,5 +1,6 @@ conf.yml *.out +*.log */**/main main dist diff --git a/app/main.go b/app/main.go index 9262b3d..f3ab1ae 100644 --- a/app/main.go +++ b/app/main.go @@ -5,6 +5,7 @@ import ( "os" "strings" + "github.com/yedf/dtm/bench" "github.com/yedf/dtm/dtmcli" "github.com/yedf/dtm/dtmsvr" "github.com/yedf/dtm/examples" @@ -43,6 +44,8 @@ func main() { // quick_start 比较独立,单独作为一个例子运行,方便新人上手 examples.QsStartSvr() examples.QsFireRequest() + case "bench": + bench.StartSvr() case "dev", "dtmsvr": default: // 下面是各类的例子 diff --git a/bench/http.go b/bench/http.go new file mode 100644 index 0000000..4b780df --- /dev/null +++ b/bench/http.go @@ -0,0 +1,124 @@ +package bench + +import ( + "database/sql" + "fmt" + "strings" + "sync/atomic" + "time" + + "github.com/gin-gonic/gin" + "github.com/yedf/dtm/common" + "github.com/yedf/dtm/dtmcli" + "github.com/yedf/dtm/examples" +) + +// 启动命令:go run app/main.go qs + +// 事务参与者的服务地址 +const benchAPI = "/api/busi_bench" +const benchPort = 8083 +const total = 1000000 + +var benchBusi = fmt.Sprintf("http://localhost:%d%s", benchPort, benchAPI) + +func sdbGet() *sql.DB { + db, err := dtmcli.PooledDB(common.DtmConfig.DB) + dtmcli.FatalIfError(err) + return db +} + +func txGet() *sql.Tx { + db := sdbGet() + tx, err := db.Begin() + dtmcli.FatalIfError(err) + return tx +} + +func reloadData() { + began := time.Now() + db := sdbGet() + tables := []string{"dtm_busi.user_account", "dtm.trans_global", "dtm.trans_branch"} + for _, t := range tables { + dtmcli.DBExec(db, fmt.Sprintf("truncate %s", t)) + } + s := "insert ignore into dtm_busi.user_account(user_id, balance) values " + ss := []string{} + for i := 1; i <= total; i++ { + ss = append(ss, fmt.Sprintf("(%d, 1000000)", i)) + } + db.Exec(s + strings.Join(ss, ",")) + dtmcli.Logf("%d users inserted. used: %dms", total, time.Since(began).Milliseconds()) +} + +var uidCounter int32 = 0 +var mode string = "" + +// StartSvr 1 +func StartSvr() { + app := common.GetGinApp() + benchAddRoute(app) + dtmcli.Logf("bench listening at %d", benchPort) + reloadData() + go app.Run(fmt.Sprintf(":%d", benchPort)) + time.Sleep(100 * time.Millisecond) +} + +func qsAdjustBalance(uid int, amount int) (interface{}, error) { + if strings.Contains(mode, "empty") { + return dtmcli.MapSuccess, nil + } else { + tx := txGet() + for i := 0; i < 5; i++ { + _, err := dtmcli.DBExec(tx, "update dtm_busi.user_account set balance = balance + ? where user_id = ?", amount, uid) + dtmcli.FatalIfError(err) + } + err := tx.Commit() + dtmcli.FatalIfError(err) + } + + return dtmcli.MapSuccess, nil +} + +func benchAddRoute(app *gin.Engine) { + app.POST(benchAPI+"/TransIn", common.WrapHandler(func(c *gin.Context) (interface{}, error) { + return qsAdjustBalance(dtmcli.MustAtoi(c.Query("uid")), 1) + })) + app.POST(benchAPI+"/TransInCompensate", common.WrapHandler(func(c *gin.Context) (interface{}, error) { + return qsAdjustBalance(dtmcli.MustAtoi(c.Query("uid")), -1) + })) + app.POST(benchAPI+"/TransOut", common.WrapHandler(func(c *gin.Context) (interface{}, error) { + return qsAdjustBalance(dtmcli.MustAtoi(c.Query("uid")), -1) + })) + app.POST(benchAPI+"/TransOutCompensate", common.WrapHandler(func(c *gin.Context) (interface{}, error) { + return qsAdjustBalance(dtmcli.MustAtoi(c.Query("uid")), 30) + })) + app.Any(benchAPI+"/reloadData", common.WrapHandler(func(c *gin.Context) (interface{}, error) { + reloadData() + mode = c.Query("m") + return nil, nil + })) + app.Any(benchAPI+"/bench", common.WrapHandler(func(c *gin.Context) (interface{}, error) { + uid := (atomic.AddInt32(&uidCounter, 1)-1)%total + 1 + suid := fmt.Sprintf("%d", uid) + suid2 := fmt.Sprintf("%d", total+1-uid) + req := gin.H{} + params := fmt.Sprintf("?uid=%s", suid) + params2 := fmt.Sprintf("?uid=%s", suid2) + dtmcli.Logf("mode: %s contains dtm: %t", mode, strings.Contains(mode, "dtm")) + if strings.Contains(mode, "dtm") { + saga := dtmcli.NewSaga(examples.DtmServer, fmt.Sprintf("bench-%d", uid)). + Add(benchBusi+"/TransOut"+params, benchBusi+"/TransOutCompensate"+params, req). + Add(benchBusi+"/TransIn"+params2, benchBusi+"/TransInCompensate"+params2, req) + saga.WaitResult = true + err := saga.Submit() + dtmcli.FatalIfError(err) + } else { + _, err := dtmcli.RestyClient.R().SetBody(gin.H{}).SetQueryParam("uid", suid2).Post(benchBusi + "/TransOut") + dtmcli.FatalIfError(err) + _, err = dtmcli.RestyClient.R().SetBody(gin.H{}).SetQueryParam("uid", suid).Post(benchBusi + "/TransIn") + dtmcli.FatalIfError(err) + } + return nil, nil + })) +} diff --git a/bench/run.sh b/bench/run.sh new file mode 100644 index 0000000..cdd1129 --- /dev/null +++ b/bench/run.sh @@ -0,0 +1,15 @@ +# !/bin/bash + +# go run ../app/main.go +set -x +TIME=1 +CURRENT=10 +curl "http://127.0.0.1:8083/api/busi_bench/reloadData?m=raw_empty" && ab -t $TIME -c $CURRENT "http://127.0.0.1:8083/api/busi_bench/bench" +curl "http://127.0.0.1:8083/api/busi_bench/reloadData?m=raw_tx" && ab -t $TIME -c $CURRENT "http://127.0.0.1:8083/api/busi_bench/bench" +curl "http://127.0.0.1:8083/api/busi_bench/reloadData?m=dtm_empty" && ab -t $TIME -c $CURRENT "http://127.0.0.1:8083/api/busi_bench/bench" +curl "http://127.0.0.1:8083/api/busi_bench/reloadData?m=dtm_tx" && ab -t $TIME -c $CURRENT "http://127.0.0.1:8083/api/busi_bench/bench" + +# curl "http://127.0.0.1:8083/api/busi_bench/reloadData?m=raw_empty" && curl "http://127.0.0.1:8083/api/busi_bench/bench" +# curl "http://127.0.0.1:8083/api/busi_bench/reloadData?m=raw_tx" && curl "http://127.0.0.1:8083/api/busi_bench/bench" +# curl "http://127.0.0.1:8083/api/busi_bench/reloadData?m=dtm_empty" && curl "http://127.0.0.1:8083/api/busi_bench/bench" +# curl "http://127.0.0.1:8083/api/busi_bench/reloadData?m=dtm_tx" && curl "http://127.0.0.1:8083/api/busi_bench/bench" diff --git a/dtmsvr/api.go b/dtmsvr/api.go index 7d22b07..3c87ed1 100644 --- a/dtmsvr/api.go +++ b/dtmsvr/api.go @@ -9,19 +9,30 @@ import ( func svcSubmit(t *TransGlobal, waitResult bool) (interface{}, error) { db := dbGet() - dbt := TransFromDb(db, t.Gid) - if dbt != nil && dbt.Status != dtmcli.StatusPrepared && dbt.Status != dtmcli.StatusSubmitted { - return M{"dtm_result": dtmcli.ResultFailure, "message": fmt.Sprintf("current status %s, cannot sumbmit", dbt.Status)}, nil - } t.Status = dtmcli.StatusSubmitted - t.saveNew(db) - return t.Process(db, waitResult), nil + err := t.saveNew(db) + if err == errUniqueConflict { + dbt := TransFromDb(db, t.Gid) + if dbt.Status == dtmcli.StatusPrepared { + updates := t.setNextCron(config.TransCronInterval) + db.Must().Model(t).Where("gid=? and status=?", t.Gid, dtmcli.StatusPrepared).Select(append(updates, "status")).Updates(t) + } else if dbt.Status != dtmcli.StatusSubmitted { + return M{"dtm_result": dtmcli.ResultFailure, "message": fmt.Sprintf("current status %s, cannot sumbmit", dbt.Status)}, nil + } + } + return t.Process(db, waitResult), nil } func svcPrepare(t *TransGlobal) (interface{}, error) { t.Status = dtmcli.StatusPrepared - t.saveNew(dbGet()) + err := t.saveNew(dbGet()) + if err == errUniqueConflict { + dbt := TransFromDb(dbGet(), t.Gid) + if dbt.Status != dtmcli.StatusPrepared { + return M{"dtm_result": dtmcli.ResultFailure, "message": fmt.Sprintf("current status %s, cannot prepare", dbt.Status)}, nil + } + } return dtmcli.MapSuccess, nil } diff --git a/dtmsvr/api_http.go b/dtmsvr/api_http.go index 45f22bf..ec59e85 100644 --- a/dtmsvr/api_http.go +++ b/dtmsvr/api_http.go @@ -25,10 +25,7 @@ func newGid(c *gin.Context) (interface{}, error) { } func prepare(c *gin.Context) (interface{}, error) { - t := TransFromContext(c) - t.Status = dtmcli.StatusPrepared - t.saveNew(dbGet()) - return dtmcli.MapSuccess, nil + return svcPrepare(TransFromContext(c)) } func submit(c *gin.Context) (interface{}, error) { diff --git a/dtmsvr/trans.go b/dtmsvr/trans.go index d2cf575..624b373 100644 --- a/dtmsvr/trans.go +++ b/dtmsvr/trans.go @@ -18,6 +18,8 @@ import ( "gorm.io/gorm/clause" ) +var errUniqueConflict = errors.New("unique key conflict error") + // TransGlobal global transaction type TransGlobal struct { common.ModelBase @@ -219,29 +221,27 @@ func (t *TransGlobal) execBranch(db *common.DB, branch *TransBranch) { } } -func (t *TransGlobal) saveNew(db *common.DB) { - err := db.Transaction(func(db1 *gorm.DB) error { +func (t *TransGlobal) saveNew(db *common.DB) error { + return db.Transaction(func(db1 *gorm.DB) error { db := &common.DB{DB: db1} - updates := t.setNextCron(config.TransCronInterval) + t.setNextCron(config.TransCronInterval) writeTransLog(t.Gid, "create trans", t.Status, "", t.Data) dbr := db.Must().Clauses(clause.OnConflict{ DoNothing: true, }).Create(t) - if dbr.RowsAffected > 0 { // 如果这个是新事务,保存所有的分支 - branches := t.getProcessor().GenBranches() - if len(branches) > 0 { - writeTransLog(t.Gid, "save branches", t.Status, "", dtmcli.MustMarshalString(branches)) - checkLocalhost(branches) - db.Must().Clauses(clause.OnConflict{ - DoNothing: true, - }).Create(&branches) - } - } else if dbr.RowsAffected == 0 && t.Status == dtmcli.StatusSubmitted { // 如果数据库已经存放了prepared的事务,则修改状态 - dbr = db.Must().Model(t).Where("gid=? and status=?", t.Gid, dtmcli.StatusPrepared).Select(append(updates, "status")).Updates(t) + if dbr.RowsAffected <= 0 { // 如果这个不是新事务,返回错误 + return errUniqueConflict + } + branches := t.getProcessor().GenBranches() + if len(branches) > 0 { + writeTransLog(t.Gid, "save branches", t.Status, "", dtmcli.MustMarshalString(branches)) + checkLocalhost(branches) + db.Must().Clauses(clause.OnConflict{ + DoNothing: true, + }).Create(&branches) } return nil }) - e2p(err) } // TransFromContext TransFromContext diff --git a/test/saga_test.go b/test/saga_test.go index 91c8d13..2081bbb 100644 --- a/test/saga_test.go +++ b/test/saga_test.go @@ -10,7 +10,6 @@ import ( ) func TestSaga(t *testing.T) { - sagaNormal(t) sagaCommittedPending(t) sagaRollback(t) @@ -23,6 +22,8 @@ func sagaNormal(t *testing.T) { assert.Equal(t, []string{dtmcli.StatusPrepared, dtmcli.StatusSucceed, dtmcli.StatusPrepared, dtmcli.StatusSucceed}, getBranchesStatus(saga.Gid)) assert.Equal(t, dtmcli.StatusSucceed, getTransStatus(saga.Gid)) transQuery(t, saga.Gid) + err := saga.Submit() // 第二次提交 + assert.Error(t, err) } func sagaCommittedPending(t *testing.T) {