From ea2d80707123c934993d84d75911fab483b09fe4 Mon Sep 17 00:00:00 2001 From: zhangwenjian Date: Sun, 27 Sep 2026 20:45:45 +0800 Subject: [PATCH] =?UTF-8?q?fix=F0=9F=90=9B:=20give=20queries=20the=20reque?= =?UTF-8?q?st's=20context,=20not=20the=20pooled=20gin.Context?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit WithContextDb and the five CRUD actions passed the *gin.Context itself to GORM as the query's context. database/sql watches the context of a query that returns rows inside a transaction from a goroutine of its own, which can still read it after the handler returns - and gin hands that Context to the next request as soon as the handler returns. GORM wraps every write in a transaction, and an INSERT on SQLite or PostgreSQL returns the new key as a row, so any write raced with the next request's reset. The race detector reports it on the second of two sequential requests. They now pass c.Request.Context(), which belongs to one request only and is cancelled when its client goes away. --- common/actions/create.go | 2 +- common/actions/delete.go | 2 +- common/actions/index.go | 2 +- common/actions/update.go | 2 +- common/actions/view.go | 2 +- common/middleware/db.go | 2 +- common/middleware/db_test.go | 60 ++++++++++++++++++++++++++++++++++++ 7 files changed, 66 insertions(+), 6 deletions(-) create mode 100644 common/middleware/db_test.go 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) + } + } +}