Skip to content

Commit 05eb259

Browse files
fix(generator): Refactor model file generation (#85)
Co-authored-by: Takafumi SEKIGUCHI <takkyuuplayer@gmail.com>
1 parent 4262e4c commit 05eb259

11 files changed

Lines changed: 434 additions & 133 deletions

README.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -129,7 +129,7 @@ func GenerateModels() {
129129
And results are mostly like this:
130130

131131
```go
132-
// This file is generated by exql. DO NOT edit.
132+
// Code generated by exql. DO NOT EDIT.
133133
package model
134134

135135
type Users struct {

generator.go

Lines changed: 49 additions & 88 deletions
Original file line numberDiff line numberDiff line change
@@ -1,17 +1,13 @@
11
package exql
22

33
import (
4-
"bytes"
54
"database/sql"
65
"fmt"
76
"go/format"
7+
"log"
88
"os"
99
"path/filepath"
10-
"strconv"
11-
"strings"
12-
"text/template"
13-
14-
"github.com/iancoleman/strcase"
10+
"regexp"
1511
)
1612

1713
type Generator interface {
@@ -24,19 +20,16 @@ type GenerateOptions struct {
2420
OutDir string
2521
Package string
2622
Exclude []string
23+
// FileNameMap maps table names to output file names.
24+
// Values must match [A-Za-z0-9_-]+.go.
25+
FileNameMap map[string]string
2726
}
2827

29-
type templateData struct {
30-
Imports string
31-
Model string
32-
ModelLower string
33-
M string
34-
Package string
35-
Fields string
36-
UpdaterFields string
37-
ScannedFields string
38-
TableName string
39-
TableNameGoLiteral string
28+
var safeModelFileNamePattern = regexp.MustCompile(`^[A-Za-z0-9_-]+\.go$`)
29+
30+
type modelFileOutput struct {
31+
path string
32+
source []byte
4033
}
4134

4235
func NewGenerator(db *sql.DB) Generator {
@@ -80,95 +73,63 @@ func (d *generator) Generate(opts *GenerateOptions) error {
8073
if err := rows.Err(); err != nil {
8174
return err
8275
}
76+
var outputs []*modelFileOutput
77+
seenPaths := map[string]string{}
8378
for _, table := range tables {
84-
if err := d.generateModelFile(table, opts); err != nil {
79+
output, err := d.generateModelFile(table, opts)
80+
if err != nil {
81+
return err
82+
}
83+
if prevTable, ok := seenPaths[output.path]; ok {
84+
return fmt.Errorf("duplicate generated model file %q for tables %q and %q", output.path, prevTable, table)
85+
}
86+
seenPaths[output.path] = table
87+
outputs = append(outputs, output)
88+
}
89+
for _, output := range outputs {
90+
if err := writeModelFile(output); err != nil {
8591
return err
8692
}
8793
}
8894
return nil
8995
}
9096

91-
func (d *generator) generateModelFile(tableName string, opt *GenerateOptions) error {
92-
tmpl := template.Must(template.New("model").Parse(modelTemplate))
97+
func (d *generator) generateModelFile(tableName string, opt *GenerateOptions) (*modelFileOutput, error) {
9398
p := NewParser()
9499
table, err := p.ParseTable(d.db, tableName)
95100
if err != nil {
96-
return err
97-
}
98-
var imports []string
99-
100-
if table.HasJsonField() {
101-
imports = append(imports, `import "encoding/json"`)
102-
}
103-
if table.HasTimeField() {
104-
imports = append(imports, `import "time"`)
101+
return nil, err
105102
}
106-
if table.HasNullField() {
107-
imports = append(imports, `import "github.com/loilo-inc/exql/v3/null"`)
103+
modelFile, err := table.GenerateModelFile(opt.Package)
104+
if err != nil {
105+
return nil, err
108106
}
109-
fields := strings.Builder{}
110-
updateFields := strings.Builder{}
111-
scannedFields := strings.Builder{}
112-
for i, col := range table.Columns {
113-
scannedFields.WriteString(fmt.Sprintf(
114-
"\t&%s.%s,", table.TableName[0:1], col.Field()),
115-
)
116-
fields.WriteString(fmt.Sprintf("\t%s", col.Field()))
117-
updateFields.WriteString(fmt.Sprintf("\t%s", col.UpdateField()))
118-
if i < len(table.Columns)-1 {
119-
scannedFields.WriteString("\n")
120-
fields.WriteString("\n")
121-
updateFields.WriteString("\n")
107+
outFileName := modelFile.Name
108+
if mappedFileName, ok := opt.FileNameMap[tableName]; ok {
109+
if err := validateMappedModelFileName(mappedFileName); err != nil {
110+
return nil, err
122111
}
112+
outFileName = mappedFileName
123113
}
124-
data := &templateData{
125-
Imports: strings.Join(imports, "\n"),
126-
Model: strcase.ToCamel(table.TableName),
127-
ModelLower: strcase.ToLowerCamel(table.TableName),
128-
M: table.TableName[0:1],
129-
UpdaterFields: updateFields.String(),
130-
Package: opt.Package,
131-
Fields: fields.String(),
132-
TableName: tableName,
133-
TableNameGoLiteral: strconv.Quote(tableName),
134-
ScannedFields: scannedFields.String(),
135-
}
136-
outFile := filepath.Join(
137-
opt.OutDir,
138-
fmt.Sprintf("%s.go", strcase.ToSnake(table.TableName)),
139-
)
140-
var buf = &bytes.Buffer{}
141-
if err := tmpl.Execute(buf, data); err != nil {
142-
return err
143-
}
144-
if fmted, err := format.Source(buf.Bytes()); err != nil {
114+
return &modelFileOutput{
115+
path: filepath.Join(opt.OutDir, outFileName),
116+
source: modelFile.Source,
117+
}, nil
118+
}
119+
120+
func writeModelFile(output *modelFileOutput) error {
121+
if fmted, err := format.Source(output.source); err != nil {
145122
return err
146-
} else if err := os.WriteFile(outFile, fmted, 0640); err != nil {
123+
} else if err := os.WriteFile(output.path, fmted, 0640); err != nil {
147124
return err
148125
}
126+
log.Printf("generated file: %s", output.path)
149127
return nil
150128
}
151129

152-
const modelTemplate = `// This file is generated by exql. DO NOT edit.
153-
package {{.Package}}
154-
155-
{{.Imports}}
156-
157-
type {{.Model}} struct {
158-
{{.Fields}}
159-
}
160-
161-
func ({{.M}} *{{.Model}}) TableName() string {
162-
return {{.Model}}TableName
163-
}
164-
165-
type Update{{.Model}} struct {
166-
{{.UpdaterFields}}
167-
}
168-
169-
func ({{.M}} *Update{{.Model}}) UpdateTableName() string {
170-
return {{.Model}}TableName
130+
func validateMappedModelFileName(name string) error {
131+
if !safeModelFileNamePattern.MatchString(name) {
132+
return fmt.Errorf("invalid model file name %q: must match [A-Za-z0-9_-]+.go", name)
133+
}
134+
return nil
171135
}
172-
173-
const {{.Model}}TableName = {{.TableNameGoLiteral}}
174-
`

generator_test.go

Lines changed: 141 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,9 @@
11
package exql
22

33
import (
4+
"bytes"
45
"fmt"
6+
"log"
57
"os"
68
"path/filepath"
79
"regexp"
@@ -136,3 +138,142 @@ func TestGenerator_Generate_EscapesTableNameGoLiteral(t *testing.T) {
136138
}
137139
assert.NoError(t, mock.ExpectationsWereMet())
138140
}
141+
142+
func TestGenerator_Generate_UsesFileNameMap(t *testing.T) {
143+
mockDb, mock, err := sqlmock.New()
144+
assert.NoError(t, err)
145+
defer mockDb.Close()
146+
147+
table := "evil/foo"
148+
mock.ExpectQuery(`show tables`).WillReturnRows(sqlmock.NewRows([]string{"tables"}).AddRow(table))
149+
mock.ExpectQuery(regexp.QuoteMeta(fmt.Sprintf("show columns from `%s`", table))).WillReturnRows(
150+
sqlmock.NewRows([]string{"Field", "Type", "Null", "Key", "Default", "Extra"}).
151+
AddRow("id", "int(11)", "NO", "PRI", nil, ""),
152+
)
153+
154+
dir := t.TempDir()
155+
var logBuf bytes.Buffer
156+
oldLogOutput := log.Writer()
157+
oldLogFlags := log.Flags()
158+
log.SetOutput(&logBuf)
159+
log.SetFlags(0)
160+
t.Cleanup(func() {
161+
log.SetOutput(oldLogOutput)
162+
log.SetFlags(oldLogFlags)
163+
})
164+
165+
err = NewGenerator(mockDb).Generate(&GenerateOptions{
166+
OutDir: dir,
167+
Package: "dist",
168+
FileNameMap: map[string]string{
169+
table: "evil.go",
170+
},
171+
})
172+
assert.NoError(t, err)
173+
174+
outFile := filepath.Join(dir, "evil.go")
175+
content, err := os.ReadFile(outFile)
176+
assert.NoError(t, err)
177+
assert.Contains(t, string(content), "package dist")
178+
assert.Contains(t, string(content), `"evil/foo"`)
179+
assert.Contains(t, logBuf.String(), outFile)
180+
assert.NoError(t, mock.ExpectationsWereMet())
181+
}
182+
183+
func TestGenerator_Generate_ReturnsErrorForDuplicateFileNames(t *testing.T) {
184+
mockDb, mock, err := sqlmock.New()
185+
assert.NoError(t, err)
186+
defer mockDb.Close()
187+
188+
mock.ExpectQuery(`show tables`).WillReturnRows(
189+
sqlmock.NewRows([]string{"tables"}).
190+
AddRow("users").
191+
AddRow("user_groups"),
192+
)
193+
mock.ExpectQuery("show columns from `users`").WillReturnRows(
194+
sqlmock.NewRows([]string{"Field", "Type", "Null", "Key", "Default", "Extra"}).
195+
AddRow("id", "int(11)", "NO", "PRI", nil, ""),
196+
)
197+
mock.ExpectQuery("show columns from `user_groups`").WillReturnRows(
198+
sqlmock.NewRows([]string{"Field", "Type", "Null", "Key", "Default", "Extra"}).
199+
AddRow("id", "int(11)", "NO", "PRI", nil, ""),
200+
)
201+
202+
dir := t.TempDir()
203+
err = NewGenerator(mockDb).Generate(&GenerateOptions{
204+
OutDir: dir,
205+
Package: "dist",
206+
FileNameMap: map[string]string{
207+
"users": "models.go",
208+
"user_groups": "models.go",
209+
},
210+
})
211+
assert.ErrorContains(t, err, "duplicate generated model file")
212+
assert.ErrorContains(t, err, "users")
213+
assert.ErrorContains(t, err, "user_groups")
214+
entries, readErr := os.ReadDir(dir)
215+
assert.NoError(t, readErr)
216+
assert.Empty(t, entries)
217+
assert.NoError(t, mock.ExpectationsWereMet())
218+
}
219+
220+
func TestGenerator_Generate_PropagatesWriteError(t *testing.T) {
221+
mockDb, mock, err := sqlmock.New()
222+
assert.NoError(t, err)
223+
defer mockDb.Close()
224+
225+
table := "users"
226+
mock.ExpectQuery(`show tables`).WillReturnRows(sqlmock.NewRows([]string{"tables"}).AddRow(table))
227+
mock.ExpectQuery("show columns from `users`").WillReturnRows(
228+
sqlmock.NewRows([]string{"Field", "Type", "Null", "Key", "Default", "Extra"}).
229+
AddRow("id", "int(11)", "NO", "PRI", nil, ""),
230+
)
231+
232+
dir := t.TempDir()
233+
err = os.Mkdir(filepath.Join(dir, "users.go"), 0750)
234+
assert.NoError(t, err)
235+
236+
err = NewGenerator(mockDb).Generate(&GenerateOptions{
237+
OutDir: dir,
238+
Package: "dist",
239+
FileNameMap: map[string]string{
240+
table: "users.go",
241+
},
242+
})
243+
assert.Error(t, err)
244+
assert.NoError(t, mock.ExpectationsWereMet())
245+
}
246+
247+
func TestValidateMappedModelFileName(t *testing.T) {
248+
for _, name := range []string{
249+
"users.go",
250+
"user_groups.go",
251+
"user-groups_1.go",
252+
"Users1.go",
253+
} {
254+
t.Run("valid "+name, func(t *testing.T) {
255+
assert.NoError(t, validateMappedModelFileName(name))
256+
})
257+
}
258+
259+
for _, name := range []string{
260+
"",
261+
".",
262+
"..",
263+
".go",
264+
"users",
265+
"users.go.txt",
266+
"user.name.go",
267+
"user name.go",
268+
"user;name.go",
269+
"ユーザー.go",
270+
"../users.go",
271+
"/tmp/users.go",
272+
"nested/users.go",
273+
`nested\users.go`,
274+
} {
275+
t.Run("invalid "+name, func(t *testing.T) {
276+
assert.Error(t, validateMappedModelFileName(name))
277+
})
278+
}
279+
}

model/fields.go

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

model/group_users.go

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

model/user_groups.go

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

model/user_login_histories.go

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

model/users.go

Lines changed: 1 addition & 1 deletion
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

0 commit comments

Comments
 (0)