Merge pull request #952 from go-admin-team/fix/gen-swallowed-errors

Code generator: stop swallowing errors when saving and generating
This commit is contained in:
wenjianzhang
2026-09-26 14:31:28 +08:00
committed by GitHub
4 changed files with 482 additions and 112 deletions
+131 -100
View File
@@ -5,6 +5,8 @@ import (
"fmt"
"go-admin/app/admin/service"
"go-admin/app/admin/service/dto"
"os"
"path/filepath"
"strconv"
"strings"
"text/template"
@@ -122,7 +124,12 @@ func (e Gen) Preview(c *gin.Context) {
return
}
tab, _ := table.Get(db, false)
tab, err := table.Get(db, false)
if err != nil {
log.Errorf("get table error, %s", err.Error())
e.Error(500, err, fmt.Sprintf("读取表配置失败!错误详情:%s", err.Error()))
return
}
// MLTBName (table_name with underscores turned to dashes) is a gorm:"-"
// field - table.Get never fills it in, so every template that reads it
// (the .vue/.ts import paths, e.g. "@/api/{PackageName}/{MLTBName}")
@@ -138,35 +145,23 @@ func (e Gen) Preview(c *gin.Context) {
// R2: infer a width for any column the config page left at colWidth's 0
// sentinel, before vue.go.template reads .ColWidth - see column_width.go.
applyInferredColumnWidths(tab.Columns)
var b1 bytes.Buffer
err = t1.Execute(&b1, tab)
var b2 bytes.Buffer
err = t2.Execute(&b2, tab)
var b3 bytes.Buffer
err = t3.Execute(&b3, tab)
var b4 bytes.Buffer
err = t4.Execute(&b4, tab)
var b5 bytes.Buffer
err = t5.Execute(&b5, tab)
var b6 bytes.Buffer
err = t6.Execute(&b6, tab)
var b7 bytes.Buffer
err = t7.Execute(&b7, tab)
var b8 bytes.Buffer
err = t8.Execute(&b8, tab)
var b9 bytes.Buffer
err = t9.Execute(&b9, tab)
out, err := renderAll(tab, t1, t2, t3, t4, t5, t6, t7, t8, t9)
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("模版渲染失败!错误详情:%s", err.Error()))
return
}
mp := make(map[string]interface{})
mp["template/model.go.template"] = b1.String()
mp["template/api.go.template"] = b2.String()
mp["template/api.ts.template"] = b3.String()
mp["template/vue.go.template"] = b4.String()
mp["template/router.go.template"] = b5.String()
mp["template/dto.go.template"] = b6.String()
mp["template/service.go.template"] = b7.String()
mp["template/lang-zh.go.template"] = b8.String()
mp["template/lang-en.go.template"] = b9.String()
mp["template/model.go.template"] = string(out[0])
mp["template/api.go.template"] = string(out[1])
mp["template/api.ts.template"] = string(out[2])
mp["template/vue.go.template"] = string(out[3])
mp["template/router.go.template"] = string(out[4])
mp["template/dto.go.template"] = string(out[5])
mp["template/service.go.template"] = string(out[6])
mp["template/lang-zh.go.template"] = string(out[7])
mp["template/lang-en.go.template"] = string(out[8])
e.OK(mp, "")
}
@@ -189,9 +184,16 @@ func (e Gen) GenCode(c *gin.Context) {
}
table.TableId = id
tab, _ := table.Get(db, false)
tab, err := table.Get(db, false)
if err != nil {
log.Errorf("get table error, %s", err.Error())
e.Error(500, err, fmt.Sprintf("读取表配置失败!错误详情:%s", err.Error()))
return
}
e.NOActionsGen(c, tab)
if !e.NOActionsGen(c, tab) {
return
}
e.OK("", "Code generated successfully!")
}
@@ -215,19 +217,29 @@ func (e Gen) GenApiToFile(c *gin.Context) {
}
table.TableId = id
tab, _ := table.Get(db, false)
e.genApiToFile(c, tab)
tab, err := table.Get(db, false)
if err != nil {
log.Errorf("get table error, %s", err.Error())
e.Error(500, err, fmt.Sprintf("读取表配置失败!错误详情:%s", err.Error()))
return
}
if !e.genApiToFile(c, tab) {
return
}
e.OK("", "Code generated successfully!")
}
func (e Gen) NOActionsGen(c *gin.Context, tab tools.SysTables) {
// NOActionsGen renders every template for tab and writes the files, and
// reports whether it did. On false it has already written the error response,
// and nothing has been written to disk if a template failed to render.
func (e Gen) NOActionsGen(c *gin.Context, tab tools.SysTables) bool {
e.Context = c
log := e.GetLogger()
tab.MLTBName = strings.Replace(tab.TBName, "_", "-", -1)
if err := requireSinglePrimaryKey(tab); err != nil {
e.Error(500, err, err.Error())
return
return false
}
// R2: see the matching call and comment in Preview above.
applyInferredColumnWidths(tab.Columns)
@@ -243,125 +255,104 @@ func (e Gen) NOActionsGen(c *gin.Context, tab tools.SysTables) {
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("model模版读取失败!错误详情:%s", err.Error()))
return
return false
}
t2, err := template.ParseFiles(basePath + "no_actions/apis.go.template")
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("api模版读取失败!错误详情:%s", err.Error()))
return
return false
}
t3, err := template.ParseFiles(routerFile)
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("路由模版失败!错误详情:%s", err.Error()))
return
return false
}
t4, err := template.ParseFiles(basePath + "ts.go.template")
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("ts模版解析失败!错误详情:%s", err.Error()))
return
return false
}
t5, err := template.ParseFiles(basePath + "vue.go.template")
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("vue模版解析失败!错误详情:%s", err.Error()))
return
return false
}
t6, err := template.ParseFiles(basePath + "dto.go.template")
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("dto模版解析失败失败!错误详情:%s", err.Error()))
return
return false
}
t7, err := template.ParseFiles(basePath + "no_actions/service.go.template")
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("service模版失败!错误详情:%s", err.Error()))
return
return false
}
// t8/t9 back F3/F9 (PRD 010): see the matching comment in Preview above.
t8, err := parseGenTemplate(basePath + "lang-zh.go.template")
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("zh语言包模版解析失败!错误详情:%s", err.Error()))
return
return false
}
t9, err := parseGenTemplate(basePath + "lang-en.go.template")
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("en语言包模版解析失败!错误详情:%s", err.Error()))
return
return false
}
_ = pkg.PathCreate("./app/" + tab.PackageName + "/apis/")
_ = pkg.PathCreate("./app/" + tab.PackageName + "/models/")
_ = pkg.PathCreate("./app/" + tab.PackageName + "/router/")
_ = pkg.PathCreate("./app/" + tab.PackageName + "/service/dto/")
_ = pkg.PathCreate(config.GenConfig.FrontPath + "/api/" + tab.PackageName + "/")
err = pkg.PathCreate(config.GenConfig.FrontPath + "/views/" + tab.PackageName + "/" + tab.MLTBName + "/")
out, err := renderAll(tab, t1, t2, t3, t4, t5, t6, t7, t8, t9)
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("views目录创建失败!错误详情:%s", err.Error()))
return
e.Error(500, err, fmt.Sprintf("模版渲染失败!错误详情:%s", err.Error()))
return false
}
// gen/{PackageName}/ nests under each locale so go-admin-ui's
// gen-namespace.ts (`./*/*.ts` glob, one level under gen/) picks the file
// up - a flat gen/{BusinessName}.ts would let two tables in different
// packages silently overwrite each other's translations, since
// BusinessName only has a pattern check, no uniqueness check.
err = pkg.PathCreate(config.GenConfig.FrontPath + "/lang/zh-CN/gen/" + tab.PackageName + "/")
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("zh语言包目录创建失败!错误详情:%s", err.Error()))
return
files := []struct {
path string
content []byte
}{
{"./app/" + tab.PackageName + "/models/" + tab.TBName + ".go", out[0]},
{"./app/" + tab.PackageName + "/apis/" + tab.TBName + ".go", out[1]},
{"./app/" + tab.PackageName + "/router/" + tab.TBName + ".go", out[2]},
{config.GenConfig.FrontPath + "/api/" + tab.PackageName + "/" + tab.MLTBName + ".ts", out[3]},
{config.GenConfig.FrontPath + "/views/" + tab.PackageName + "/" + tab.MLTBName + "/index.vue", out[4]},
{"./app/" + tab.PackageName + "/service/dto/" + tab.TBName + ".go", out[5]},
{"./app/" + tab.PackageName + "/service/" + tab.TBName + ".go", out[6]},
{config.GenConfig.FrontPath + "/lang/zh-CN/gen/" + tab.PackageName + "/" + tab.BusinessName + ".ts", out[7]},
{config.GenConfig.FrontPath + "/lang/en-US/gen/" + tab.PackageName + "/" + tab.BusinessName + ".ts", out[8]},
}
err = pkg.PathCreate(config.GenConfig.FrontPath + "/lang/en-US/gen/" + tab.PackageName + "/")
if err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("en语言包目录创建失败!错误详情:%s", err.Error()))
return
for _, f := range files {
if err := writeGenerated(f.path, f.content); err != nil {
log.Error(err)
e.Error(500, err, fmt.Sprintf("生成文件写入失败!错误详情:%s", err.Error()))
return false
}
}
var b1 bytes.Buffer
err = t1.Execute(&b1, tab)
var b2 bytes.Buffer
err = t2.Execute(&b2, tab)
var b3 bytes.Buffer
err = t3.Execute(&b3, tab)
var b4 bytes.Buffer
err = t4.Execute(&b4, tab)
var b5 bytes.Buffer
err = t5.Execute(&b5, tab)
var b6 bytes.Buffer
err = t6.Execute(&b6, tab)
var b7 bytes.Buffer
err = t7.Execute(&b7, tab)
var b8 bytes.Buffer
err = t8.Execute(&b8, tab)
var b9 bytes.Buffer
err = t9.Execute(&b9, tab)
pkg.FileCreate(b1, "./app/"+tab.PackageName+"/models/"+tab.TBName+".go")
pkg.FileCreate(b2, "./app/"+tab.PackageName+"/apis/"+tab.TBName+".go")
pkg.FileCreate(b3, "./app/"+tab.PackageName+"/router/"+tab.TBName+".go")
pkg.FileCreate(b4, config.GenConfig.FrontPath+"/api/"+tab.PackageName+"/"+tab.MLTBName+".ts")
pkg.FileCreate(b5, config.GenConfig.FrontPath+"/views/"+tab.PackageName+"/"+tab.MLTBName+"/index.vue")
pkg.FileCreate(b6, "./app/"+tab.PackageName+"/service/dto/"+tab.TBName+".go")
pkg.FileCreate(b7, "./app/"+tab.PackageName+"/service/"+tab.TBName+".go")
pkg.FileCreate(b8, config.GenConfig.FrontPath+"/lang/zh-CN/gen/"+tab.PackageName+"/"+tab.BusinessName+".ts")
pkg.FileCreate(b9, config.GenConfig.FrontPath+"/lang/en-US/gen/"+tab.PackageName+"/"+tab.BusinessName+".ts")
return true
}
func (e Gen) genApiToFile(c *gin.Context, tab tools.SysTables) {
// genApiToFile writes the migration that seeds tab's menu and APIs, and
// reports whether it did; on false the error response is already written.
func (e Gen) genApiToFile(c *gin.Context, tab tools.SysTables) bool {
err := e.MakeContext(c).
MakeOrm().
Errors
if err != nil {
e.Logger.Error(err)
e.Error(500, err, err.Error())
return
return false
}
basePath := "template/"
@@ -370,17 +361,24 @@ func (e Gen) genApiToFile(c *gin.Context, tab tools.SysTables) {
if err != nil {
e.Logger.Error(err)
e.Error(500, err, fmt.Sprintf("数据迁移模版解析失败!错误详情:%s", err.Error()))
return
return false
}
i := strconv.FormatInt(time.Now().UnixNano()/1e6, 10)
var b1 bytes.Buffer
err = t1.Execute(&b1, struct {
out, err := renderAll(struct {
tools.SysTables
GenerateTime string
}{tab, i})
pkg.FileCreate(b1, "./cmd/migrate/migration/version-local/"+i+"_migrate.go")
}{tab, i}, t1)
if err != nil {
e.Logger.Error(err)
e.Error(500, err, fmt.Sprintf("数据迁移模版渲染失败!错误详情:%s", err.Error()))
return false
}
if err := writeGenerated("./cmd/migrate/migration/version-local/"+i+"_migrate.go", out[0]); err != nil {
e.Logger.Error(err)
e.Error(500, err, fmt.Sprintf("数据迁移文件写入失败!错误详情:%s", err.Error()))
return false
}
return true
}
func (e Gen) GenMenuAndApi(c *gin.Context) {
@@ -404,7 +402,12 @@ func (e Gen) GenMenuAndApi(c *gin.Context) {
}
table.TableId = id
tab, _ := table.Get(e.Orm, true)
tab, err := table.Get(e.Orm, true)
if err != nil {
e.Logger.Errorf("get table error, %s", err.Error())
e.Error(500, err, fmt.Sprintf("读取表配置失败!错误详情:%s", err.Error()))
return
}
tab.MLTBName = strings.Replace(tab.TBName, "_", "-", -1)
Mmenu := dto.SysMenuInsertReq{}
@@ -533,3 +536,31 @@ func requireSinglePrimaryKey(tab tools.SysTables) error {
return fmt.Errorf("表 %s 是联合主键(%s),代码生成需要恰好一个主键列", tab.TBName, strings.Join(keys, ", "))
}
}
// renderAll executes each template against data and returns the outputs in
// the same order, stopping at the first template that fails. The caller
// writes nothing until every template has rendered: a template that fails
// part-way leaves a truncated file behind it, which compiles as nothing and
// at a glance looks like the real thing.
func renderAll(data any, tpls ...*template.Template) ([][]byte, error) {
out := make([][]byte, len(tpls))
for i, t := range tpls {
var b bytes.Buffer
if err := t.Execute(&b, data); err != nil {
return nil, err
}
out[i] = b.Bytes()
}
return out, nil
}
// writeGenerated writes one generated file, creating its directory first.
// It stands in for pkg.FileCreate, which returns no error at all and, when
// the file cannot be created, closes a nil file and ends the process with
// log.Fatalln - one unwritable path took the whole server down with it.
func writeGenerated(path string, content []byte) error {
if err := os.MkdirAll(filepath.Dir(path), os.ModePerm); err != nil {
return err
}
return os.WriteFile(path, content, 0o644)
}
+217
View File
@@ -0,0 +1,217 @@
package tools
import (
"encoding/json"
"io/fs"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"testing"
"github.com/gin-gonic/gin"
"github.com/glebarez/sqlite"
"github.com/go-admin-team/go-admin-core/v2/sdk/config"
"gorm.io/gorm"
"gorm.io/gorm/logger"
"go-admin/app/other/models/tools"
)
type genBody struct {
Code int `json:"code"`
Msg string `json:"msg"`
}
// runGen calls h and decodes every JSON body it wrote, so a handler that
// answers twice - an error, then a success - shows up as two bodies.
func runGen(t *testing.T, db *gorm.DB, h func(*gin.Context), params gin.Params) []genBody {
t.Helper()
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodGet, "/", nil)
c.Params = params
if db != nil {
c.Set("db", db)
}
h(c)
var bodies []genBody
dec := json.NewDecoder(strings.NewReader(w.Body.String()))
for dec.More() {
var b genBody
if err := dec.Decode(&b); err != nil {
t.Fatalf("decoding %q: %v", w.Body.String(), err)
}
bodies = append(bodies, b)
}
return bodies
}
// genWorkspace is a working directory laid out the way the generator expects
// to find it: the real templates under template/, output written below it.
// Any template can be swapped for one that fails when executed.
func genWorkspace(t *testing.T, broken ...string) string {
t.Helper()
gin.SetMode(gin.TestMode)
root, err := filepath.Abs("../../../..")
if err != nil {
t.Fatal(err)
}
dir := t.TempDir()
src := filepath.Join(root, "template")
err = filepath.WalkDir(src, func(p string, d fs.DirEntry, err error) error {
if err != nil {
return err
}
dst := filepath.Join(dir, "template", strings.TrimPrefix(p, src))
if d.IsDir() {
return os.MkdirAll(dst, 0o755)
}
b, err := os.ReadFile(p)
if err != nil {
return err
}
return os.WriteFile(dst, b, 0o644)
})
if err != nil {
t.Fatal(err)
}
for _, name := range broken {
// Parses, then fails on execution: SysTables has no such field.
if err := os.WriteFile(filepath.Join(dir, "template", name), []byte("{{.NoSuchField}}"), 0o644); err != nil {
t.Fatal(err)
}
}
front := config.GenConfig.FrontPath
t.Cleanup(func() { config.GenConfig.FrontPath = front })
config.GenConfig.FrontPath = filepath.Join(dir, "ui")
t.Chdir(dir)
return dir
}
func genTable() tools.SysTables {
return tools.SysTables{
TBName: "gen_err", ClassName: "GenErr", BusinessName: "genErr", PackageName: "admin",
ModuleName: "gen-err", PkColumn: "id", PkGoField: "Id", PkJsonField: "id",
Columns: []tools.SysColumns{
{ColumnName: "id", GoField: "Id", GoType: "int", JsonField: "id", Pk: true, IsPk: "1"},
{ColumnName: "name", GoField: "Name", GoType: "string", JsonField: "name", IsInsert: "1"},
},
}
}
func writtenFiles(t *testing.T, dir string) []string {
t.Helper()
var files []string
for _, top := range []string{"app", "ui"} {
_ = filepath.WalkDir(filepath.Join(dir, top), func(p string, d fs.DirEntry, err error) error {
if err == nil && !d.IsDir() {
files = append(files, p)
}
return nil
})
}
return files
}
// A template that fails on execution used to leave its file empty and every
// other file written, with the request reporting success.
func TestNOActionsGenWritesNothingWhenATemplateFails(t *testing.T) {
dir := genWorkspace(t, "v4/no_actions/service.go.template")
var ok bool
bodies := runGen(t, nil, func(c *gin.Context) { ok = Gen{}.NOActionsGen(c, genTable()) }, nil)
if ok || len(bodies) != 1 || bodies[0].Code != 500 {
t.Fatalf("ok=%v, responses %+v; want false and one 500", ok, bodies)
}
if !strings.Contains(bodies[0].Msg, "service.go.template") {
t.Errorf("message %q does not name the template that failed", bodies[0].Msg)
}
if files := writtenFiles(t, dir); len(files) > 0 {
t.Errorf("wrote %d file(s) although a template failed: %v", len(files), files)
}
}
// pkg.FileCreate, which this used to write through, ends the process with
// log.Fatalln when it cannot create the file. Reaching the assertions at all
// is half of this test.
func TestNOActionsGenReportsAFileItCannotWrite(t *testing.T) {
dir := genWorkspace(t)
blocker := filepath.Join(dir, "app", "admin", "models")
if err := os.MkdirAll(filepath.Dir(blocker), 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(blocker, nil, 0o644); err != nil { // a file where a directory must go
t.Fatal(err)
}
var ok bool
bodies := runGen(t, nil, func(c *gin.Context) { ok = Gen{}.NOActionsGen(c, genTable()) }, nil)
if ok || len(bodies) != 1 || bodies[0].Code != 500 {
t.Fatalf("ok=%v, responses %+v; want false and one 500", ok, bodies)
}
// Not a template failing first for some other reason.
if !strings.Contains(bodies[0].Msg, "生成文件写入失败") {
t.Errorf("message %q is not the write failure", bodies[0].Msg)
}
}
func genDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatal(err)
}
if err := db.AutoMigrate(new(tools.SysTables), new(tools.SysColumns)); err != nil {
t.Fatal(err)
}
return db
}
// GenCode answered the request itself after NOActionsGen had already answered
// it with an error, so the client received the error followed by
// "Code generated successfully".
func TestGenCodeAnswersOnceWhenGenerationFails(t *testing.T) {
genWorkspace(t, "v4/no_actions/service.go.template")
db := genDB(t)
table := genTable()
created, err := table.Create(db)
if err != nil {
t.Fatal(err)
}
bodies := runGen(t, db, Gen{}.GenCode, gin.Params{{Key: "tableId", Value: itoa(created.TableId)}})
if len(bodies) != 1 || bodies[0].Code != 500 {
t.Errorf("responses %+v; want exactly one, a 500", bodies)
}
}
// Every handler that reads a table's configuration by id used to discard the
// lookup's error and carry on with an empty table: GenCode and Preview then
// refused it as having no primary key, and GenMenuAndApi went on to seed
// menus and APIs for it.
func TestGeneratorHandlersReportAMissingTable(t *testing.T) {
genWorkspace(t)
db := genDB(t)
missing := gin.Params{{Key: "tableId", Value: "987654"}}
for name, h := range map[string]func(*gin.Context){
"GenCode": Gen{}.GenCode,
"Preview": Gen{}.Preview,
"GenApiToFile": Gen{}.GenApiToFile,
"GenMenuAndApi": Gen{}.GenMenuAndApi,
} {
t.Run(name, func(t *testing.T) {
bodies := runGen(t, db, h, missing)
if len(bodies) != 1 || bodies[0].Code != 500 || !strings.Contains(bodies[0].Msg, "读取表配置失败") {
t.Errorf("responses %+v; want one 500 saying the table could not be read", bodies)
}
})
}
}
+32 -12
View File
@@ -1,6 +1,7 @@
package tools
import (
"fmt"
common "go-admin/common/models"
"strings"
@@ -132,22 +133,29 @@ func (e *SysTables) GetTree(tx *gorm.DB) ([]SysTables, error) {
return doc, nil
}
// Create inserts the table and every one of its columns, or none of them. A
// column that failed to insert used to be dropped without a word, leaving a
// table the config page lists with columns missing and code generated from
// fewer columns than the table has.
func (e *SysTables) Create(tx *gorm.DB) (SysTables, error) {
var doc SysTables
e.CreateBy = 0
result := tx.Table("sys_tables").Create(&e)
if result.Error != nil {
err := result.Error
err := tx.Transaction(func(tx *gorm.DB) error {
if err := tx.Table("sys_tables").Create(&e).Error; err != nil {
return err
}
for i := range e.Columns {
e.Columns[i].TableId = e.TableId
if _, err := e.Columns[i].Create(tx); err != nil {
return fmt.Errorf("column %s: %w", e.Columns[i].ColumnName, err)
}
}
return nil
})
if err != nil {
return doc, err
}
doc = *e
for i := 0; i < len(e.Columns); i++ {
e.Columns[i].TableId = doc.TableId
_, _ = e.Columns[i].Create(tx)
}
return doc, nil
return *e, nil
}
func (e *SysTables) Update(tx *gorm.DB) (update SysTables, err error) {
@@ -158,6 +166,16 @@ func (e *SysTables) Update(tx *gorm.DB) (update SysTables, err error) {
//参数1:是要修改的数据
//参数2:是修改的数据
e.UpdateBy = 0
// One transaction for the table and its columns, as in Create: a column
// whose write failed used to be skipped while the request reported
// success, so the config page showed a save that had not happened.
err = tx.Transaction(func(tx *gorm.DB) error {
return e.updateWithColumns(tx)
})
return
}
func (e *SysTables) updateWithColumns(tx *gorm.DB) (err error) {
if err = tx.Table("sys_tables").Where("table_id = ?", e.TableId).Updates(&e).Error; err != nil {
return
}
@@ -199,7 +217,9 @@ func (e *SysTables) Update(tx *gorm.DB) (update SysTables, err error) {
}
}
}
_, _ = e.Columns[i].Update(tx)
if _, err = e.Columns[i].Update(tx); err != nil {
return fmt.Errorf("column %s: %w", e.Columns[i].ColumnName, err)
}
}
return
}
@@ -0,0 +1,102 @@
package tools
import (
"errors"
"testing"
"github.com/glebarez/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// A named, shared in-memory database: with a bare ":memory:" every pooled
// connection opens its own empty database, and a row written inside a
// transaction would be invisible to the query that checks for it.
func openAtomicDB(t *testing.T) *gorm.DB {
t.Helper()
db, err := gorm.Open(sqlite.Open("file:"+t.Name()+"?mode=memory&cache=shared"), &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
})
if err != nil {
t.Fatalf("open sqlite: %v", err)
}
if err := db.AutoMigrate(new(SysTables), new(SysColumns)); err != nil {
t.Fatalf("migrate: %v", err)
}
return db
}
// rejectNthColumnWrite makes the n-th write of a sys_columns row fail, the way
// it does when the database refuses it.
func rejectNthColumnWrite(register func(name string, fn func(*gorm.DB)) error, n int) error {
seen := 0
return register("test:reject_column", func(tx *gorm.DB) {
if tx.Statement.Table != "sys_columns" {
return
}
seen++
if seen == n {
_ = tx.AddError(errors.New("column rejected"))
}
})
}
func threeColumns() []SysColumns {
return []SysColumns{{ColumnName: "id"}, {ColumnName: "name"}, {ColumnName: "note"}}
}
func TestSysTablesCreateWritesAllColumnsOrNone(t *testing.T) {
db := openAtomicDB(t)
if err := rejectNthColumnWrite(db.Callback().Create().Before("gorm:create").Register, 2); err != nil {
t.Fatal(err)
}
table := SysTables{TBName: "orders", Columns: threeColumns()}
if _, err := table.Create(db); err == nil {
t.Fatal("Create reported success although a column was rejected")
}
var tables, columns int64
db.Model(new(SysTables)).Where("table_name = ?", "orders").Count(&tables)
db.Model(new(SysColumns)).Count(&columns)
if tables != 0 || columns != 0 {
t.Errorf("left %d table row(s) and %d column row(s) behind; want none", tables, columns)
}
}
func TestSysTablesUpdateWritesAllColumnsOrNone(t *testing.T) {
db := openAtomicDB(t)
table := SysTables{TBName: "orders", TableComment: "before", Columns: threeColumns()}
created, err := table.Create(db)
if err != nil {
t.Fatal(err)
}
if err := rejectNthColumnWrite(db.Callback().Update().Before("gorm:update").Register, 2); err != nil {
t.Fatal(err)
}
edited, err := (&SysTables{TableId: created.TableId}).Get(db, false)
if err != nil {
t.Fatal(err)
}
edited.TableComment = "after"
for i := range edited.Columns {
edited.Columns[i].ColumnComment = "after"
}
if _, err := edited.Update(db); err == nil {
t.Fatal("Update reported success although a column was rejected")
}
stored, err := (&SysTables{TableId: created.TableId}).Get(db, false)
if err != nil {
t.Fatal(err)
}
if stored.TableComment != "before" {
t.Errorf("table comment = %q; the table row was kept although a column failed", stored.TableComment)
}
for _, c := range stored.Columns {
if c.ColumnComment != "" {
t.Errorf("column %s comment = %q; a column write was kept although another failed", c.ColumnName, c.ColumnComment)
}
}
}