diff --git a/common/actions/create.go b/common/actions/create.go index f667fc97..9fe75bd7 100644 --- a/common/actions/create.go +++ b/common/actions/create.go @@ -37,7 +37,7 @@ func CreateAction(control dto.Control) gin.HandlerFunc { return } object.SetCreateBy(user.GetUserId(c)) - err = db.WithContext(c).Create(object).Error + err = db.WithContext(c.Request.Context()).Create(object).Error if err != nil { log.Errorf("Create error: %s", err) response.Error(c, 500, err, "创建失败") diff --git a/common/actions/delete.go b/common/actions/delete.go index b8e992b6..530ec997 100644 --- a/common/actions/delete.go +++ b/common/actions/delete.go @@ -43,7 +43,7 @@ func DeleteAction(control dto.Control) gin.HandlerFunc { //数据权限检查 p := GetPermissionFromContext(c) - db = db.WithContext(c).Scopes( + db = db.WithContext(c.Request.Context()).Scopes( Permission(object.TableName(), p), ).Where(req.GetId()).Delete(object) if err = db.Error; err != nil { diff --git a/common/actions/index.go b/common/actions/index.go index 17758c65..eb9afb40 100644 --- a/common/actions/index.go +++ b/common/actions/index.go @@ -39,7 +39,7 @@ func IndexAction(m models.ActiveRecord, d dto.Index, f func() interface{}) gin.H //数据权限检查 p := GetPermissionFromContext(c) - err = db.WithContext(c).Model(object). + err = db.WithContext(c.Request.Context()).Model(object). Scopes( dto.MakeCondition(req.GetNeedSearch()), dto.Paginate(req.GetPageSize(), req.GetPageIndex()), diff --git a/common/actions/update.go b/common/actions/update.go index 2c0afde6..3ce96d07 100644 --- a/common/actions/update.go +++ b/common/actions/update.go @@ -41,7 +41,7 @@ func UpdateAction(control dto.Control) gin.HandlerFunc { //数据权限检查 p := GetPermissionFromContext(c) - db = db.WithContext(c).Scopes( + db = db.WithContext(c.Request.Context()).Scopes( Permission(object.TableName(), p), ).Where(req.GetId()).Updates(object) if err = db.Error; err != nil { diff --git a/common/actions/view.go b/common/actions/view.go index 03d26821..3f3a0f2f 100644 --- a/common/actions/view.go +++ b/common/actions/view.go @@ -48,7 +48,7 @@ func ViewAction(control dto.Control, f func() interface{}) gin.HandlerFunc { //数据权限检查 p := GetPermissionFromContext(c) - err = db.Model(object).WithContext(c).Scopes( + err = db.Model(object).WithContext(c.Request.Context()).Scopes( Permission(object.TableName(), p), ).Where(req.GetId()).First(rsp).Error diff --git a/common/middleware/db.go b/common/middleware/db.go index 905eb04e..06d2bcc7 100644 --- a/common/middleware/db.go +++ b/common/middleware/db.go @@ -6,6 +6,6 @@ import ( ) func WithContextDb(c *gin.Context) { - c.Set("db", sdk.Runtime.GetDbByTenant(c.Request.Host).WithContext(c)) + c.Set("db", sdk.Runtime.GetDbByTenant(c.Request.Host).WithContext(c.Request.Context())) c.Next() } diff --git a/common/middleware/db_test.go b/common/middleware/db_test.go new file mode 100644 index 00000000..c150fb76 --- /dev/null +++ b/common/middleware/db_test.go @@ -0,0 +1,60 @@ +package middleware + +import ( + "net/http" + "net/http/httptest" + "testing" + + "github.com/gin-gonic/gin" + "github.com/glebarez/sqlite" + "github.com/go-admin-team/go-admin-core/v2/sdk" + "gorm.io/gorm" +) + +// database/sql watches the context of a query that returns rows inside a +// transaction from a goroutine of its own, which can still be reading it +// after the handler has returned. GORM wraps every write in a transaction, +// and an INSERT on SQLite or PostgreSQL returns the new key as a row. gin +// reuses its Context for the next request as soon as the handler returns, so +// a query given the gin.Context races with that reuse; given the request's +// own context, it does not. Sequential requests are enough for the reuse, +// and -race, which make test runs with, reports it. +type probeRow struct { + Id int `gorm:"primaryKey;autoIncrement"` +} + +func (probeRow) TableName() string { return "with_context_db_probe" } + +func TestWithContextDbDoesNotHandQueriesThePooledContext(t *testing.T) { + const host = "with-context-db.test" + db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{}) + if err != nil { + t.Fatal(err) + } + if err := db.AutoMigrate(&probeRow{}); err != nil { + t.Fatal(err) + } + previous := sdk.Runtime.GetDbByTenant(host) + sdk.Runtime.SetDbByTenant(host, db) + t.Cleanup(func() { sdk.Runtime.SetDbByTenant(host, previous) }) + + gin.SetMode(gin.TestMode) + r := gin.New() + r.Use(WithContextDb) + r.GET("/", func(c *gin.Context) { + if err := c.MustGet("db").(*gorm.DB).Create(&probeRow{}).Error; err != nil { + c.Status(http.StatusInternalServerError) + return + } + c.Status(http.StatusOK) + }) + for i := 0; i < 20; i++ { + w := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.Host = host + r.ServeHTTP(w, req) + if w.Code != http.StatusOK { + t.Fatalf("request %d answered %d", i, w.Code) + } + } +}