22 Commits

Author SHA1 Message Date
83a7087a3b TMP: codetables
Some checks failed
CI / build-docker (push) Successful in 4s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Failing after 6s
2026-09-16 18:11:13 -07:00
4f3865deaa codegen: generate getters for compound indexes
Some checks failed
CI / build-docker (push) Successful in 4s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Failing after 6s
2026-09-15 21:38:33 -07:00
3f2db59b3d refactor (codegen): reduce redundancy, make 'get item(s) by blah' generators more general 2026-09-15 18:19:36 -07:00
89b2fd7c3c codegen: generate getters for non-unique indexes 2026-09-15 17:51:52 -07:00
a1b06a3ef8 codegen: use must.Be helper in Save function 2026-09-08 18:27:03 -07:00
49d2a7748f codegen: make generated tests support 'without rowid' tables 2026-09-08 17:26:47 -07:00
ae036d15f2 codegen: make test fixture factory function intelligent about what fields and values it assigns 2026-09-08 17:01:04 -07:00
b33127db6c codegen: make generated code use 'must' instead of 'flowutils' 2026-09-08 15:35:17 -07:00
f542f45630 refactor: migrate 'flowutils' to 'must' helpers 2026-09-08 15:21:29 -07:00
cf989f7433 codegen: add error printout for when codegen re-parse fails to facilitate debugging 2026-09-08 14:51:31 -07:00
a101c7531b codegen: make modelSQLFields const use camel case instead of lowercase
All checks were successful
CI / build-docker (push) Successful in 3s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Successful in 18s
2026-08-18 11:12:57 -07:00
9c31266f59 schema: split 'column', 'index' and 'schema' into separate files from 'table' 2026-08-18 11:12:47 -07:00
616304c7dd ci: add support for manually triggering the build
All checks were successful
CI / build-docker (push) Successful in 5s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Successful in 4m14s
2026-07-12 14:05:05 -07:00
17fc8a68f6 codegen: wrap foreign key checks for nullable FKs in "if a.Val != 0 { ... }" 2026-07-12 13:34:21 -07:00
d572745613 codegen: don't skip created_at auto-timestamp for 'without rowid' tables 2026-07-12 13:20:35 -07:00
9bfb31798c codegen: don't auto-timestamp overwrite provided timestamps if there are ones 2026-07-07 12:15:51 -07:00
eafeb658bd codegen: fix invalid SQL query being generated for GetItemBy with multiple params 2026-06-24 13:59:13 -07:00
dbf14e23b6 codegen: fix fk checking lambda producing non-lint-passing code
All checks were successful
CI / build-docker (push) Successful in 4s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Successful in 2m42s
2026-05-29 17:34:40 -07:00
ad1782c73d codegen: use "require.NoError(...)" when saving objects that return errors 2026-05-29 15:36:36 -07:00
ccd7e32cbf codegen: test file now uses deep.Equal instead of comparing one field 2026-05-29 14:06:09 -07:00
ed4ade1956 codegen: fix escaped double-quote inside test strings 2026-05-29 13:44:13 -07:00
9a11f3986c codegen: improve the "TestFkChecking" generated function
All checks were successful
CI / build-docker (push) Successful in 6s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Successful in 3m22s
2026-04-24 14:05:48 -07:00
14 changed files with 680 additions and 365 deletions

View File

@@ -1,6 +1,6 @@
name: CI
on: [push]
on: [push, workflow_dispatch]
jobs:
# These steps build the `gas` docker image.

View File

@@ -60,8 +60,6 @@ linters:
- -ST1000 # Re-enable this once we have docstrings
- -ST1003 # I like snake_case
- -ST1013 # HTTP status codes are shorter and more readable than names
dot-import-whitelist:
- "git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"
exclusions:
generated: lax # Don't lint generated files
paths:

View File

@@ -6,11 +6,12 @@ import (
"go/ast"
"go/token"
"os"
"slices"
"github.com/spf13/cobra"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/codegen/modelgenerate"
. "git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/schema"
)
@@ -23,22 +24,22 @@ var generate_model = &cobra.Command{
Args: cobra.ExactArgs(1),
RunE: func(cmd *cobra.Command, args []string) error {
path := Must(cmd.Flags().GetString("schema"))
modname := Must(cmd.Flags().GetString("modname"))
path := must.Get(cmd.Flags().GetString("schema"))
modname := must.Get(cmd.Flags().GetString("modname"))
sql, err := os.ReadFile(path)
if err != nil {
return fmt.Errorf("reading path %s: %w", path, err)
}
db := schema.InitDB(string(sql))
schema := schema.SchemaFromDB(db)
table, isOk := schema.Tables[args[0]]
sch := schema.SchemaFromDB(db)
table, isOk := sch.Tables[args[0]]
if !isOk {
return ErrNoSuchTable
}
if Must(cmd.Flags().GetBool("test")) {
file2 := modelgenerate.GenerateModelTestAST(table, schema, modname)
PanicIf(modelgenerate.FprintWithComments(os.Stdout, file2))
if must.Get(cmd.Flags().GetBool("test")) {
file2 := modelgenerate.GenerateModelTestAST(table, sch, modname)
must.Do(modelgenerate.FprintWithComments(os.Stdout, file2))
} else {
decls := []ast.Decl{
&ast.GenDecl{
@@ -51,10 +52,9 @@ var generate_model = &cobra.Command{
Path: &ast.BasicLit{Kind: token.STRING, Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/db"`},
},
&ast.ImportSpec{
Name: ast.NewIdent("."),
Path: &ast.BasicLit{
Kind: token.STRING,
Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"`,
Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"`,
},
},
},
@@ -78,13 +78,23 @@ var generate_model = &cobra.Command{
modelgenerate.GenerateGetItemByIDFunc(table),
)
}
for _, index := range schema.Indexes {
for _, index := range sch.Indexes {
if index.TableName != table.TableName {
// Skip indexes on other tables
continue
}
if index.IsUnique && len(index.Columns) == 1 {
decls = append(decls, modelgenerate.GenerateGetItemByUniqColFunc(table, table.GetColumnByName(index.Columns[0])))
if slices.Contains(index.Columns, "") {
// Skip expression indexes; there's no way to resolve an expression to a real column
continue
}
cols := make([]schema.Column, len(index.Columns))
for i, colName := range index.Columns {
cols[i] = table.GetColumnByName(colName)
}
if index.IsUnique {
decls = append(decls, modelgenerate.GenerateGetItemBy(table, cols))
} else {
decls = append(decls, modelgenerate.GenerateGetItemsBy(table, cols))
}
}
decls = append(decls, modelgenerate.GenerateGetAllItemsFunc(table))
@@ -94,7 +104,7 @@ var generate_model = &cobra.Command{
Decls: decls,
}
PanicIf(modelgenerate.FprintWithComments(os.Stdout, file))
must.Do(modelgenerate.FprintWithComments(os.Stdout, file))
}
return nil

View File

@@ -9,7 +9,7 @@ import (
"github.com/spf13/cobra"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/codegen"
. "git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"
)
var cmd_init = &cobra.Command{
@@ -23,11 +23,11 @@ var cmd_init = &cobra.Command{
var target string
if len(args) != 0 {
target = args[0]
PanicIf(os.MkdirAll(target, 0o755))
PanicIf(os.Chdir(target))
must.Do(os.MkdirAll(target, 0o755))
must.Do(os.Chdir(target))
} else {
// Default to current directory (".")
target = Must(os.Getwd())
target = must.Get(os.Getwd())
}
// Get all the config values
@@ -41,9 +41,9 @@ var cmd_init = &cobra.Command{
}
}
pkg_opts := codegen.PkgOpts{
ModuleName: Must(cmd.Flags().GetString("module")),
DBFilename: Must(cmd.Flags().GetString("db")),
BinaryName: Must(cmd.Flags().GetString("binary")),
ModuleName: must.Get(cmd.Flags().GetString("module")),
DBFilename: must.Get(cmd.Flags().GetString("db")),
BinaryName: must.Get(cmd.Flags().GetString("binary")),
}
if pkg_opts.ModuleName == "" {
pkg_opts.ModuleName = filepath.Base(target)

3
ops/test.sh Normal file
View File

@@ -0,0 +1,3 @@
#!/bin/sh
go test -tags fts5 ./...

View File

@@ -27,10 +27,18 @@ const (
// we don't have to do that.
var TrailingComments = map[ast.Node]string{}
// mustCall wraps a call expression in Must(...), producing AST for Must(inner).
// mustCall wraps a call expression in must.Get(...), producing AST for must.Get(inner).
func mustCall(inner ast.Expr) *ast.CallExpr {
return &ast.CallExpr{
Fun: ast.NewIdent("Must"),
Fun: &ast.SelectorExpr{X: ast.NewIdent("must"), Sel: ast.NewIdent("Get")},
Args: []ast.Expr{inner},
}
}
// doCall wraps a call expression in must.Do(...), producing AST for must.Do(inner).
func doCall(inner ast.Expr) *ast.CallExpr {
return &ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("must"), Sel: ast.NewIdent("Do")},
Args: []ast.Expr{inner},
}
}
@@ -73,7 +81,7 @@ func FprintWithComments(w io.Writer, file *ast.File) error {
fset = token.NewFileSet()
parsed, err := parser.ParseFile(fset, "", buf.Bytes(), parser.ParseComments)
if err != nil {
return fmt.Errorf("re-parsing pretty-print: %w", err)
return fmt.Errorf("re-parsing pretty-print: %w\n\n%s", err, buf.String())
}
// Parallel walk: apply TrailingComments from the side map.

View File

@@ -5,6 +5,7 @@ import (
"fmt"
"go/ast"
"go/token"
"slices"
"strings"
"github.com/jinzhu/inflection"
@@ -23,7 +24,7 @@ var (
)
func SQLFieldsConstIdent(tbl schema.Table) *ast.Ident {
return ast.NewIdent(strings.ToLower(tbl.GoTypeName) + "SQLFields")
return ast.NewIdent(strings.ToLower(tbl.GoTypeName[:1]) + tbl.GoTypeName[1:] + "SQLFields")
}
// GoTypeForColumn returns a type expression for this column.
@@ -53,20 +54,48 @@ func GoTypeForColumn(c schema.Column) ast.Expr {
}
}
func PanicIfRowsAffected(tbl schema.Table) *ast.IfStmt {
return &ast.IfStmt{
Cond: &ast.BinaryExpr{
// MustBeRowsAffected produces an AST for a `must.Be(...)` call asserting that exactly one row
// was affected by the preceding statement, e.g.:
//
// must.Be(must.Get(result.RowsAffected()) == 1, "%w: Food ID=%d", ErrNotInDB, f.ID)
//
// For "without rowid" tables, the message includes the table's primary key column(s) instead of ID.
func MustBeRowsAffected(tbl schema.Table) *ast.ExprStmt {
pkParts := []string{}
pkArgs := []ast.Expr{}
if tbl.IsWithoutRowid {
for _, col := range tbl.PrimaryKeyColumns() {
verb := "%v"
if col.Type == "integer" || col.Type == "int" || col.IsNonCodeTableForeignKey() {
verb = "%d"
}
pkParts = append(pkParts, fmt.Sprintf("%s=%s", col.Name, verb))
pkArgs = append(pkArgs, &ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent(col.GoFieldName())})
}
} else {
pkParts = append(pkParts, "ID=%d")
pkArgs = append(pkArgs, &ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("ID")})
}
msg := fmt.Sprintf("%%w: %s %s", tbl.GoTypeName, strings.Join(pkParts, ", "))
args := append([]ast.Expr{
&ast.BinaryExpr{
X: mustCall(&ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("result"), Sel: ast.NewIdent("RowsAffected")},
Args: []ast.Expr{},
}),
Op: token.NEQ,
Op: token.EQL,
Y: &ast.BasicLit{Kind: token.INT, Value: "1"},
},
Body: &ast.BlockStmt{List: []ast.Stmt{
&ast.ExprStmt{X: &ast.CallExpr{Fun: ast.NewIdent("panic"), Args: []ast.Expr{ast.NewIdent(tbl.VarName)}}},
}},
}
&ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf("%q", msg)},
ast.NewIdent("ErrNotInDB"),
}, pkArgs...)
return &ast.ExprStmt{X: &ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("must"), Sel: ast.NewIdent("Be")},
Args: args,
}}
}
// ---------------
@@ -161,9 +190,21 @@ func buildFKCheckLambda(tbl schema.Table) (*ast.AssignStmt, bool) {
structFieldName := col.GoFieldName()
structField := &ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent(structFieldName)}
ret = append(ret, func() ast.Stmt {
// Wrap nullable FKs in "if a.val != 0 { ... }"
wrap := func(input ast.Stmt) ast.Stmt {
if col.IsNullableForeignKey() {
return &ast.IfStmt{
Cond: &ast.BinaryExpr{X: structField, Op: token.NEQ, Y: &ast.BasicLit{Kind: token.INT, Value: "0"}},
Body: &ast.BlockStmt{List: []ast.Stmt{input}},
}
} else {
return input
}
}
if col.IsNonCodeTableForeignKey() {
// Real foreign key; look up referent by ID to see if it exists
ret = append(ret, &ast.IfStmt{
return wrap(&ast.IfStmt{
Init: &ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("_"), ast.NewIdent("err")},
Tok: token.DEFINE,
@@ -197,10 +238,10 @@ func buildFKCheckLambda(tbl schema.Table) (*ast.AssignStmt, bool) {
})
} else {
// Code table value. Query the table to see if it exists
ret = append(ret, &ast.IfStmt{
return wrap(&ast.IfStmt{
Init: &ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("err")},
Tok: token.ASSIGN,
Tok: token.DEFINE,
Rhs: []ast.Expr{
&ast.CallExpr{
Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("Get")},
@@ -234,6 +275,7 @@ func buildFKCheckLambda(tbl schema.Table) (*ast.AssignStmt, bool) {
},
})
}
}())
}
// final return nil
ret = append(ret, &ast.ReturnStmt{Results: []ast.Expr{ast.NewIdent("nil")}})
@@ -320,7 +362,7 @@ func GenerateSaveItemFunc(tbl schema.Table) *ast.FuncDecl {
},
}
if !hasFks {
// No foreign key checking needed; just use `Must` for brevity
// No foreign key checking needed; just use `must.Get` for brevity
return []ast.Stmt{&ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("result")},
Tok: token.DEFINE,
@@ -386,8 +428,26 @@ func GenerateSaveItemFunc(tbl schema.Table) *ast.FuncDecl {
}
if tbl.IsWithoutRowid {
if hasCreatedAt {
// Auto-timestamps: created_at. Don't overwrite existing timestamps (e.g., data import / migrations)
ret = append(ret, &ast.IfStmt{
Cond: &ast.CallExpr{Fun: &ast.SelectorExpr{
X: &ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("CreatedAt")},
Sel: ast.NewIdent("IsZero"),
}},
Body: &ast.BlockStmt{
List: []ast.Stmt{
&ast.AssignStmt{
Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("CreatedAt")}},
Tok: token.ASSIGN,
Rhs: []ast.Expr{&ast.CallExpr{Fun: ast.NewIdent("TimestampNow"), Args: []ast.Expr{}}},
},
},
},
})
}
ret = append(ret, namedExecStmt(upsertStmt)...)
ret = append(ret, PanicIfRowsAffected(tbl))
ret = append(ret, MustBeRowsAffected(tbl))
} else {
// if item.ID == 0 {...} else {...}
ret = append(ret, &ast.IfStmt{
@@ -403,10 +463,21 @@ func GenerateSaveItemFunc(tbl schema.Table) *ast.FuncDecl {
ret1 := []ast.Stmt{Comment("Do create")}
if hasCreatedAt {
// Auto-timestamps: created_at
ret1 = append(ret1, &ast.AssignStmt{
ret1 = append(ret1, &ast.IfStmt{
// Don't overwrite existing timestamps. This is useful for various reasons, e.g., data import / migrations
Cond: &ast.CallExpr{Fun: &ast.SelectorExpr{
X: &ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("CreatedAt")},
Sel: ast.NewIdent("IsZero"),
}},
Body: &ast.BlockStmt{
List: []ast.Stmt{
&ast.AssignStmt{
Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("CreatedAt")}},
Tok: token.ASSIGN,
Rhs: []ast.Expr{&ast.CallExpr{Fun: ast.NewIdent("TimestampNow"), Args: []ast.Expr{}}},
},
},
},
})
}
return append(ret1, namedExecStmt(insertStmt)...)
@@ -430,7 +501,7 @@ func GenerateSaveItemFunc(tbl schema.Table) *ast.FuncDecl {
[]ast.Stmt{Comment("Do update")},
append(
namedExecStmt(updateStmt),
PanicIfRowsAffected(tbl),
MustBeRowsAffected(tbl),
)...,
),
},
@@ -469,15 +540,50 @@ func getByIDFuncName(tblname string) string {
return "Get" + schema.TypenameFromTablename(tblname) + "ByID"
}
// paramNamesFor picks parameter identifiers for a `Get...By()` function's columns.
// A single column always uses the short GoVarName() form (e.g. "name"). For multiple
// columns, GoVarName() is preferred for readability, but if two or more of the given
// columns would produce the same short name (e.g. two foreign keys that both abbreviate
// to "uID"), the longer, unambiguous LongGoVarName() form is used for all of them instead,
// since a collision there would produce invalid (duplicate-parameter) Go code.
func paramNamesFor(cols []schema.Column) []string {
names := make([]string, len(cols))
if len(cols) == 1 {
names[0] = cols[0].GoVarName()
return names
}
hasConflict := false
for i, c := range cols {
names[i] = c.GoVarName()
if slices.Contains(names[:i], names[i]) {
hasConflict = true
break
}
}
if !hasConflict {
return names
}
for i, c := range cols {
names[i] = c.LongGoVarName()
}
return names
}
// GenerateGetItemBy produces an AST for a `GetXyzByCol1AndCol2...()` function that returns
// the single item matching an exact-match lookup over the given columns (or ErrNotInDB).
// Used for unique indexes (single- or multi-column) and for "without rowid" tables' primary keys.
func GenerateGetItemBy(tbl schema.Table, cols []schema.Column) *ast.FuncDecl {
paramNames := paramNamesFor(cols)
colNames := []string{}
funcNameSuffix := []string{}
funcParams := &ast.FieldList{List: []*ast.Field{}}
sqlParams := []ast.Expr{}
for _, col := range cols {
funcParam := ast.NewIdent(col.LongGoVarName())
for i, col := range cols {
funcParam := ast.NewIdent(paramNames[i])
funcParams.List = append(funcParams.List, &ast.Field{Names: []*ast.Ident{funcParam}, Type: GoTypeForColumn(col)})
colNames = append(colNames, fmt.Sprintf("%s = :%s", col.Name, col.Name))
colNames = append(colNames, fmt.Sprintf("%s = ?", col.Name))
funcNameSuffix = append(funcNameSuffix, col.GoFieldName())
sqlParams = append(sqlParams, funcParam)
}
@@ -489,7 +595,7 @@ func GenerateGetItemBy(tbl schema.Table, cols []schema.Column) *ast.FuncDecl {
Y: SQLFieldsConstIdent(tbl),
},
Op: token.ADD,
Y: &ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf("`\n\t from %s\n\t where %s = ?\n\t`", tbl.TableName, strings.Join(colNames, " and "))},
Y: &ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf("`\n\t from %s\n\t where %s\n\t`", tbl.TableName, strings.Join(colNames, " and "))},
}
return &ast.FuncDecl{
@@ -515,7 +621,8 @@ func GenerateGetItemBy(tbl schema.Table, cols []schema.Column) *ast.FuncDecl {
&ast.IfStmt{
Cond: &ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("errors"), Sel: ast.NewIdent("Is")},
Args: []ast.Expr{ast.NewIdent("err"), &ast.SelectorExpr{X: ast.NewIdent("sql"), Sel: ast.NewIdent("ErrNoRows")}}},
Args: []ast.Expr{ast.NewIdent("err"), &ast.SelectorExpr{X: ast.NewIdent("sql"), Sel: ast.NewIdent("ErrNoRows")}},
},
Body: &ast.BlockStmt{List: []ast.Stmt{&ast.ReturnStmt{Results: []ast.Expr{&ast.CompositeLit{Type: ast.NewIdent(tbl.GoTypeName)}, ast.NewIdent("ErrNotInDB")}}}},
},
&ast.ReturnStmt{},
@@ -564,10 +671,33 @@ func GenerateGetItemByIDFunc(tbl schema.Table) *ast.FuncDecl {
}
}
// GenerateGetItemByUniqColFunc produces an AST for the `GetXyzByID()` function.
// E.g., a table with `table.TypeName = "foods"` will produce a "GetFoodByID()" function.
// GenerateGetItemByUniqColFunc produces an AST for the `GetXyzByCol()` function, for a unique
// index on a single column.
// E.g., a table with `table.TypeName = "foods"` will produce a "GetFoodByName()" function.
func GenerateGetItemByUniqColFunc(tbl schema.Table, col schema.Column) *ast.FuncDecl {
// Use the xyzSQLFields constant in the select query
return GenerateGetItemBy(tbl, []schema.Column{col})
}
// GenerateGetItemsBy produces an AST for a `GetXyzsByCol1AndCol2...()` function that returns
// all items matching an exact-match lookup over the given columns. Used for non-unique
// indexes (single- or multi-column).
// E.g., a table with `table.TableName = "foods"` and a non-unique index on "category" will
// produce a "GetFoodsByCategory(category string) []Food" function.
func GenerateGetItemsBy(tbl schema.Table, cols []schema.Column) *ast.FuncDecl {
paramNames := paramNamesFor(cols)
colNames := []string{}
funcNameSuffix := []string{}
funcParams := &ast.FieldList{List: []*ast.Field{}}
sqlParams := []ast.Expr{}
for i, col := range cols {
funcParam := ast.NewIdent(paramNames[i])
funcParams.List = append(funcParams.List, &ast.Field{Names: []*ast.Ident{funcParam}, Type: GoTypeForColumn(col)})
colNames = append(colNames, fmt.Sprintf("%s = ?", col.Name))
funcNameSuffix = append(funcNameSuffix, col.GoFieldName())
sqlParams = append(sqlParams, funcParam)
}
selectExpr := &ast.BinaryExpr{
X: &ast.BinaryExpr{
X: &ast.BasicLit{Kind: token.STRING, Value: "`\n\t select `"},
@@ -575,40 +705,40 @@ func GenerateGetItemByUniqColFunc(tbl schema.Table, col schema.Column) *ast.Func
Y: SQLFieldsConstIdent(tbl),
},
Op: token.ADD,
Y: &ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf("`\n\t from %s\n\t where %s = ?\n\t`", tbl.TableName, col.Name)},
Y: &ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf("`\n\t from %s\n\t where %s\n\t`", tbl.TableName, strings.Join(colNames, " and "))},
}
param := ast.NewIdent(col.GoVarName())
selectCall := doCall(&ast.CallExpr{
Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("Select")},
Args: append([]ast.Expr{&ast.UnaryExpr{Op: token.AND, X: ast.NewIdent("ret")}, selectExpr}, sqlParams...),
})
return &ast.FuncDecl{
Recv: dbRecv,
Name: ast.NewIdent("Get" + schema.TypenameFromTablename(tbl.TableName) + "By" + col.GoFieldName()),
Name: ast.NewIdent("Get" + inflection.Plural(schema.TypenameFromTablename(tbl.TableName)) + "By" + strings.Join(funcNameSuffix, "And")),
Type: &ast.FuncType{
Params: &ast.FieldList{List: []*ast.Field{
{Names: []*ast.Ident{param}, Type: GoTypeForColumn(col)},
}},
Params: funcParams,
Results: &ast.FieldList{List: []*ast.Field{
{Names: []*ast.Ident{ast.NewIdent("ret")}, Type: ast.NewIdent(tbl.GoTypeName)},
{Names: []*ast.Ident{ast.NewIdent("err")}, Type: ast.NewIdent("error")},
{Names: []*ast.Ident{ast.NewIdent("ret")}, Type: &ast.ArrayType{Elt: ast.NewIdent(tbl.GoTypeName)}},
}},
},
Body: &ast.BlockStmt{
List: []ast.Stmt{
&ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("err")},
Tok: token.ASSIGN,
Rhs: []ast.Expr{&ast.CallExpr{Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("Get")}, Args: []ast.Expr{&ast.UnaryExpr{Op: token.AND, X: ast.NewIdent("ret")}, selectExpr, param}}},
},
&ast.IfStmt{
Cond: &ast.CallExpr{Fun: &ast.SelectorExpr{X: ast.NewIdent("errors"), Sel: ast.NewIdent("Is")}, Args: []ast.Expr{ast.NewIdent("err"), &ast.SelectorExpr{X: ast.NewIdent("sql"), Sel: ast.NewIdent("ErrNoRows")}}},
Body: &ast.BlockStmt{List: []ast.Stmt{&ast.ReturnStmt{Results: []ast.Expr{&ast.CompositeLit{Type: ast.NewIdent(tbl.GoTypeName)}, ast.NewIdent("ErrNotInDB")}}}},
},
&ast.ExprStmt{X: selectCall},
&ast.ReturnStmt{},
},
},
}
}
// GenerateGetItemsByColFunc produces an AST for the `GetXyzsByCol()` function, for a non-unique
// index on a single column.
// E.g., a table with `table.TableName = "foods"` and a non-unique index on "category" will
// produce a "GetFoodsByCategory(category string) []Food" function.
func GenerateGetItemsByColFunc(tbl schema.Table, col schema.Column) *ast.FuncDecl {
return GenerateGetItemsBy(tbl, []schema.Column{col})
}
// GenerateGetAllItemsFunc produces an AST for the `GetAllXyzs()` function.
// E.g., a table with `table.TypeName = "foods"` will produce a "GetAllFoods()" function.
func GenerateGetAllItemsFunc(tbl schema.Table) *ast.FuncDecl {
@@ -617,10 +747,7 @@ func GenerateGetAllItemsFunc(tbl schema.Table) *ast.FuncDecl {
{Names: []*ast.Ident{ast.NewIdent("ret")}, Type: &ast.ArrayType{Elt: ast.NewIdent(tbl.GoTypeName)}},
}}
selectCall := &ast.CallExpr{
Fun: ast.NewIdent("PanicIf"),
Args: []ast.Expr{
&ast.CallExpr{
selectCall := doCall(&ast.CallExpr{
Fun: &ast.SelectorExpr{
X: dbDB,
Sel: ast.NewIdent("Select"),
@@ -637,9 +764,7 @@ func GenerateGetAllItemsFunc(tbl schema.Table) *ast.FuncDecl {
Y: &ast.BasicLit{Kind: token.STRING, Value: "` from " + tbl.TableName + "`"},
},
},
},
},
}
})
funcBody := &ast.BlockStmt{
List: []ast.Stmt{
@@ -682,7 +807,7 @@ func GenerateDeleteItemFunc(tbl schema.Table) *ast.FuncDecl {
},
})},
},
PanicIfRowsAffected(tbl),
MustBeRowsAffected(tbl),
},
}

View File

@@ -4,6 +4,8 @@ import (
"fmt"
"go/ast"
"go/token"
"slices"
"strings"
"github.com/jinzhu/inflection"
@@ -11,6 +13,80 @@ import (
"git.offline-twitter.com/offline-labs/gas-stack/pkg/textutils"
)
// SampleValue returns a deterministic value of the appropriate type for the given column.
func SampleValue(c pkgschema.Column, offset int) ast.Expr {
switch c.Type {
case "integer", "int":
if strings.HasPrefix(c.Name, "is_") || strings.HasPrefix(c.Name, "has_") {
// Boolean case
if offset%2 == 0 {
return ast.NewIdent("false")
}
return ast.NewIdent("true")
} else if strings.HasSuffix(c.Name, "_at") {
// Timestamp case
return &ast.CallExpr{
Fun: ast.NewIdent("TimestampFromUnix"),
Args: []ast.Expr{
&ast.BasicLit{Kind: token.INT, Value: fmt.Sprintf("%d", 10000+offset)},
},
}
} else {
// Regular integer case
return &ast.BasicLit{Kind: token.INT, Value: fmt.Sprintf("%d", 10+offset)}
}
case "text":
if offset == 0 {
return &ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf("%q", "asdf")}
}
return &ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf("%q", fmt.Sprintf("asdf%d", offset))}
case "real":
return &ast.BasicLit{Kind: token.FLOAT, Value: fmt.Sprintf("%.2f", 1.23+float64(offset))}
case "blob":
return &ast.CompositeLit{
Type: &ast.ArrayType{
Elt: ast.NewIdent("byte"),
},
Elts: func() []ast.Expr {
ret := []ast.Expr{
&ast.BasicLit{Kind: token.INT, Value: "72"}, // 'H'
&ast.BasicLit{Kind: token.INT, Value: "105"}, // 'i'
&ast.BasicLit{Kind: token.INT, Value: "33"}, // '!'
}
for range offset {
ret = append(ret,
&ast.BasicLit{Kind: token.INT, Value: "33"}, // '!'
)
}
return ret
}(),
}
default:
panic("Unrecognized sqlite column type: " + c.Type)
}
}
// UpdateTestFields returns the columns for which the test generator should assign its own
// sample values, in the MakeXyz() factory and in the create/update test: the same columns
// that GenerateSaveItemFunc's "update" branch writes to, minus foreign keys (which can't be
// given plausible values here) and the auto-managed "created_at"/"updated_at" columns.
func UpdateTestFields(tbl pkgschema.Table) (ret []pkgschema.Column) {
hasCreatedAt, hasUpdatedAt := tbl.HasAutoTimestamps()
for _, c := range tbl.Columns {
if c.Name == "rowid" || c.IsPrimaryKey || c.IsForeignKey {
continue
}
if c.Name == "created_at" && hasCreatedAt {
continue
}
if c.Name == "updated_at" && hasUpdatedAt {
continue
}
ret = append(ret, c)
}
return
}
// GenerateModelTestAST produces an AST for a starter test file for a given model.
func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodName string) *ast.File {
packageName := "db"
@@ -18,6 +94,9 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
makeHelperName := ast.NewIdent("Make" + tbl.GoTypeName)
hasCreatedAt, hasUpdatedAt := tbl.HasAutoTimestamps()
updateTestFields := UpdateTestFields(tbl)
// func MakeItem() Item { return Item{} }
makeItemFunc := &ast.FuncDecl{
Name: makeHelperName,
@@ -35,24 +114,15 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
Results: []ast.Expr{
&ast.CompositeLit{
Type: ast.NewIdent(tbl.GoTypeName),
Elts: []ast.Expr{
&ast.KeyValueExpr{
Key: ast.NewIdent("Data"),
Value: &ast.CompositeLit{
Type: &ast.ArrayType{
Elt: ast.NewIdent("byte"),
},
Elts: []ast.Expr{},
},
},
&ast.KeyValueExpr{
Key: ast.NewIdent("Description"),
Value: &ast.BasicLit{
Kind: token.STRING,
Value: `""`,
},
},
},
Elts: func() (ret []ast.Expr) {
for _, c := range updateTestFields {
ret = append(ret, &ast.KeyValueExpr{
Key: ast.NewIdent(c.GoFieldName()),
Value: SampleValue(c, 0),
})
}
return
}(),
},
},
},
@@ -62,12 +132,72 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
testObj := ast.NewIdent(textutils.CamelToPascal(tbl.GoTypeName))
testObj2 := ast.NewIdent(textutils.CamelToPascal(tbl.GoTypeName) + "2")
fieldName := ast.NewIdent("Description") // TODO
description1 := `"an item"`
description2 := `"a big item"`
testDB := ast.NewIdent("TestDB")
hasCreatedAt, hasUpdatedAt := tbl.HasAutoTimestamps()
// getItemByPKCall builds a call to this table's primary-key getter (e.g. `TestDB.GetItemByID(item.ID)`,
// or `TestDB.GetItemByColAAndColB(item.ColA, item.ColB)` for "without rowid" tables with a
// compound primary key), matching whatever GenerateGetItemByIDFunc/GenerateGetItemBy generated.
getItemByPKCall := func(obj *ast.Ident) *ast.CallExpr {
if !tbl.IsWithoutRowid {
// Normal rowid table: use GetXyzByID
return &ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Get" + tbl.GoTypeName + "ByID")},
Args: []ast.Expr{&ast.SelectorExpr{X: obj, Sel: ast.NewIdent("ID")}},
}
} else {
// "Without rowid" table: use the primary key "GetItemByBlahBlah" query func
funcNameSuffix := []string{}
args := []ast.Expr{}
for _, c := range tbl.PrimaryKeyColumns() {
funcNameSuffix = append(funcNameSuffix, c.GoFieldName())
args = append(args, &ast.SelectorExpr{X: obj, Sel: ast.NewIdent(c.GoFieldName())})
}
return &ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Get" + tbl.GoTypeName + "By" + strings.Join(funcNameSuffix, "And"))},
Args: args,
}
}
}
makeDeepEqual := func(obj1 *ast.Ident, obj2 *ast.Ident) *ast.IfStmt {
return &ast.IfStmt{
Init: &ast.AssignStmt{
Lhs: []ast.Expr{
&ast.Ident{Name: "diff"},
},
Tok: token.DEFINE,
Rhs: []ast.Expr{
&ast.CallExpr{
Fun: &ast.SelectorExpr{
X: &ast.Ident{Name: "deep"},
Sel: &ast.Ident{Name: "Equal"},
},
Args: []ast.Expr{obj1, obj2},
},
},
},
Cond: &ast.BinaryExpr{
X: &ast.Ident{Name: "diff"},
Op: token.NEQ,
Y: &ast.Ident{Name: "nil"},
},
Body: &ast.BlockStmt{
List: []ast.Stmt{
&ast.ExprStmt{
X: &ast.CallExpr{
Fun: &ast.SelectorExpr{
X: &ast.Ident{Name: "t"},
Sel: &ast.Ident{Name: "Error"},
},
Args: []ast.Expr{
&ast.Ident{Name: "diff"},
},
},
},
},
},
}
}
testFuncType := &ast.FuncType{
Params: &ast.FieldList{
@@ -78,6 +208,103 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
},
}
// Generate FK Check test func first, because it also detects whether there are foreign keys
shouldIncludeTestFkCheck := false
testFkChecking := &ast.FuncDecl{
Name: ast.NewIdent("Test" + tbl.GoTypeName + "FkChecking"),
Type: testFuncType,
Body: &ast.BlockStmt{
List: func() (stmts []ast.Stmt) {
isFirst := true
for _, col := range tbl.Columns {
if !col.IsForeignKey {
continue
}
shouldIncludeTestFkCheck = true
// post := MakePost()
if !isFirst {
stmts = append(stmts, BlankLine())
}
stmts = append(stmts, []ast.Stmt{
// Comment header
Comment(fmt.Sprintf("Invalid %s", col.GoFieldName())),
// `Invalid
&ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent(tbl.VarName)},
Tok: map[bool]token.Token{true: token.DEFINE, false: token.ASSIGN}[isFirst],
Rhs: []ast.Expr{
&ast.CallExpr{
Fun: ast.NewIdent("Make" + tbl.GoTypeName),
},
},
},
// `post.QuotedPostID = 94354538969386985`
&ast.AssignStmt{
Lhs: []ast.Expr{
&ast.SelectorExpr{
X: ast.NewIdent(tbl.VarName),
Sel: ast.NewIdent(col.GoFieldName()),
},
},
Tok: token.ASSIGN,
Rhs: []ast.Expr{
&ast.BasicLit{
Kind: token.INT,
Value: "94354538969386985",
},
},
},
// `err := db.SavePost(&post)`
&ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("err")},
Tok: map[bool]token.Token{true: token.DEFINE, false: token.ASSIGN}[isFirst],
Rhs: []ast.Expr{
&ast.CallExpr{
Fun: &ast.SelectorExpr{
X: testDB,
Sel: ast.NewIdent("Save" + tbl.GoTypeName),
},
Args: []ast.Expr{
&ast.UnaryExpr{
Op: token.AND,
X: ast.NewIdent(tbl.VarName),
},
},
},
},
},
// `assertForeignKeyError(t, err, "QuotedPostID", post.QuotedPostID)`
&ast.ExprStmt{
X: &ast.CallExpr{
Fun: ast.NewIdent("AssertForeignKeyError"),
Args: []ast.Expr{
ast.NewIdent("t"),
ast.NewIdent("err"),
&ast.BasicLit{
Kind: token.STRING,
Value: fmt.Sprintf("%q", col.GoFieldName()),
},
&ast.SelectorExpr{
X: ast.NewIdent(tbl.VarName),
Sel: ast.NewIdent(col.GoFieldName()),
},
},
},
},
}...)
isFirst = false
}
return stmts
}(),
},
}
testCreateUpdateDelete := &ast.FuncDecl{
Name: ast.NewIdent("TestCreateUpdateDelete" + tbl.GoTypeName),
Type: testFuncType,
@@ -99,34 +326,45 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
Tok: token.DEFINE,
Rhs: []ast.Expr{&ast.CallExpr{Fun: makeHelperName, Args: nil}},
},
// item.Description = "an item"
&ast.AssignStmt{
// item.Description = <sample value, offset 1>
// ...one assignment per updatable field
}
for _, c := range updateTestFields {
stmts = append(stmts, &ast.AssignStmt{
Lhs: []ast.Expr{
&ast.SelectorExpr{
X: testObj,
Sel: ast.NewIdent("Description"),
Sel: ast.NewIdent(c.GoFieldName()),
},
},
Tok: token.ASSIGN,
Rhs: []ast.Expr{
&ast.BasicLit{
Kind: token.STRING,
Value: fmt.Sprintf("%q", description1),
},
},
},
// TestDB.SaveItem(&item)
&ast.ExprStmt{X: &ast.CallExpr{
Rhs: []ast.Expr{SampleValue(c, 1)},
})
}
stmts = append(stmts,
// TestDB.SaveItem(&item), possibly with error check
&ast.ExprStmt{X: func() *ast.CallExpr {
mainExpr := &ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Save" + tbl.GoTypeName)},
Args: []ast.Expr{&ast.UnaryExpr{Op: token.AND, X: testObj}},
}},
}
if shouldIncludeTestFkCheck {
// Also a check for whether the Save function returns an error
return &ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("require"), Sel: ast.NewIdent("NoError")},
Args: []ast.Expr{ast.NewIdent("t"), mainExpr},
}
}
return mainExpr
}()},
)
// require.NotZero(t, item.ID)
&ast.ExprStmt{X: &ast.CallExpr{
if !tbl.IsWithoutRowid { // non-rowid tables don't get an ID
stmts = append(stmts, &ast.ExprStmt{X: &ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("require"), Sel: ast.NewIdent("NotZero")},
Args: []ast.Expr{ast.NewIdent("t"), &ast.SelectorExpr{X: testObj, Sel: ast.NewIdent("ID")}},
}},
}})
}
// After create: assert timestamps are set
@@ -141,63 +379,48 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
BlankLine(),
Comment("Load"),
// item2 := Must(TestDB.GetItemByID(item.ID))
// item2 := must.Get(TestDB.GetItemByID(item.ID))
&ast.AssignStmt{
Lhs: []ast.Expr{testObj2},
Tok: token.DEFINE,
Rhs: []ast.Expr{mustCall(&ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Get" + tbl.GoTypeName + "ByID")},
Args: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: ast.NewIdent("ID")}},
})},
Rhs: []ast.Expr{mustCall(getItemByPKCall(testObj))},
},
// assert.Equal(t, item.Description, item2.Description)
&ast.ExprStmt{X: &ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("assert"), Sel: ast.NewIdent("Equal")},
Args: []ast.Expr{
ast.NewIdent("t"),
&ast.SelectorExpr{X: testObj, Sel: fieldName},
&ast.SelectorExpr{X: testObj2, Sel: fieldName},
},
}},
// if deep.Equal(...) {...}
makeDeepEqual(testObj, testObj2),
)
stmts = append(stmts,
BlankLine(),
Comment("Update"),
)
// item.Description = "a big item"
&ast.AssignStmt{
Lhs: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: fieldName}},
// item.Description = <sample value, offset 2>
// ...one assignment per updatable field
for _, c := range updateTestFields {
stmts = append(stmts, &ast.AssignStmt{
Lhs: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: ast.NewIdent(c.GoFieldName())}},
Tok: token.ASSIGN,
Rhs: []ast.Expr{&ast.BasicLit{Kind: token.STRING, Value: description2}},
},
Rhs: []ast.Expr{SampleValue(c, 2)},
})
}
stmts = append(stmts,
// TestDB.SaveItem(&item)
&ast.ExprStmt{X: &ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Save" + tbl.GoTypeName)},
Args: []ast.Expr{&ast.UnaryExpr{Op: token.AND, X: testObj}},
}},
// item2 = Must(TestDB.GetItemByID(item.ID))
// item2 = must.Get(TestDB.GetItemByID(item.ID))
&ast.AssignStmt{
Lhs: []ast.Expr{testObj2},
Tok: token.ASSIGN,
Rhs: []ast.Expr{mustCall(&ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Get" + tbl.GoTypeName + "ByID")},
Args: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: ast.NewIdent("ID")}},
})},
Rhs: []ast.Expr{mustCall(getItemByPKCall(testObj))},
},
// assert.Equal(t, item.Description, item2.Description)
&ast.ExprStmt{X: &ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("assert"), Sel: ast.NewIdent("Equal")},
Args: []ast.Expr{
ast.NewIdent("t"),
&ast.SelectorExpr{X: testObj, Sel: fieldName},
&ast.SelectorExpr{X: testObj2, Sel: fieldName},
},
}},
// if deep.Equal(...) {...}
makeDeepEqual(testObj, testObj2),
)
indexGets, hasIndexedGets := []ast.Stmt{
@@ -210,8 +433,21 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
// Skip indexes on other tables
continue
}
if index.IsUnique && len(index.Columns) == 1 {
col := tbl.GetColumnByName(index.Columns[0])
if slices.Contains(index.Columns, "") {
// Skip expression indexes; there's no way to resolve an expression to a real column
continue
}
cols := make([]pkgschema.Column, len(index.Columns))
nameParts := make([]string, len(index.Columns))
callArgs := make([]ast.Expr, len(index.Columns))
for i, colName := range index.Columns {
cols[i] = tbl.GetColumnByName(colName)
nameParts[i] = cols[i].GoFieldName()
callArgs[i] = &ast.SelectorExpr{X: testObj2, Sel: ast.NewIdent(cols[i].GoFieldName())}
}
funcNameSuffix := strings.Join(nameParts, "And")
if index.IsUnique {
indexGets = append(indexGets, []ast.Stmt{
// assert.Equal(t, item2, TestDB.GetItemByXYZ(...))
&ast.ExprStmt{X: &ast.CallExpr{ // TODO: what if just delete the "ExprStmt" wrapper?
@@ -221,16 +457,32 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
testObj2,
mustCall(&ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent(
"Get" + pkgschema.TypenameFromTablename(tbl.TableName) + "By" + col.GoFieldName(),
"Get" + pkgschema.TypenameFromTablename(tbl.TableName) + "By" + funcNameSuffix,
)},
Args: []ast.Expr{&ast.SelectorExpr{X: testObj2, Sel: ast.NewIdent(col.GoFieldName())}},
Args: callArgs,
}),
},
}},
}...)
// decls = append(decls, modelgenerate.GenerateGetItemByUniqColFunc(table, table.GetColumnByName(index.Columns[0])))
hasIndexedGets = true
} else {
indexGets = append(indexGets, []ast.Stmt{
// assert.Contains(t, TestDB.GetItemsByXYZ(...), item2)
&ast.ExprStmt{X: &ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("assert"), Sel: ast.NewIdent("Contains")},
Args: []ast.Expr{
ast.NewIdent("t"),
&ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent(
"Get" + inflection.Plural(pkgschema.TypenameFromTablename(tbl.TableName)) + "By" + funcNameSuffix,
)},
Args: callArgs,
},
testObj2,
},
}},
}...)
}
hasIndexedGets = true
}
if hasIndexedGets {
@@ -251,10 +503,7 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
&ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("_"), ast.NewIdent("err")},
Tok: token.DEFINE,
Rhs: []ast.Expr{&ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Get" + tbl.GoTypeName + "ByID")},
Args: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: ast.NewIdent("ID")}},
}},
Rhs: []ast.Expr{getItemByPKCall(testObj)},
},
// assert.ErrorIs(t, err, db.ErrNotInDB)
@@ -292,93 +541,6 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
},
}
shouldIncludeTestFkCheck := false
testFkChecking := &ast.FuncDecl{
Name: ast.NewIdent("Test" + tbl.GoTypeName + "FkChecking"),
Type: testFuncType,
Body: &ast.BlockStmt{
List: func() []ast.Stmt {
// post := MakePost()
stmts := []ast.Stmt{
&ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent(tbl.VarName)},
Tok: token.DEFINE,
Rhs: []ast.Expr{
&ast.CallExpr{
Fun: ast.NewIdent("Make" + tbl.GoTypeName),
},
},
},
}
shouldDefineErr := true
for _, col := range tbl.Columns {
if col.IsForeignKey {
shouldIncludeTestFkCheck = true
stmts = append(stmts, []ast.Stmt{
// post.QuotedPostID = 94354538969386985
&ast.AssignStmt{
Lhs: []ast.Expr{
&ast.SelectorExpr{
X: ast.NewIdent(tbl.VarName),
Sel: ast.NewIdent(col.GoFieldName()),
},
},
Tok: token.ASSIGN,
Rhs: []ast.Expr{
&ast.BasicLit{
Kind: token.INT,
Value: "94354538969386985",
},
},
},
// err := db.SavePost(&post)
&ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("err")},
Tok: map[bool]token.Token{true: token.DEFINE, false: token.ASSIGN}[shouldDefineErr],
Rhs: []ast.Expr{
&ast.CallExpr{
Fun: &ast.SelectorExpr{
X: testDB,
Sel: ast.NewIdent("Save" + tbl.GoTypeName),
},
Args: []ast.Expr{
&ast.UnaryExpr{
Op: token.AND,
X: ast.NewIdent(tbl.VarName),
},
},
},
},
},
// assertForeignKeyError(t, err, "QuotedPostID", post.QuotedPostID)
&ast.ExprStmt{
X: &ast.CallExpr{
Fun: ast.NewIdent("AssertForeignKeyError"),
Args: []ast.Expr{
ast.NewIdent("t"),
ast.NewIdent("err"),
&ast.BasicLit{
Kind: token.STRING,
Value: fmt.Sprintf("%q", col.GoFieldName()),
},
&ast.SelectorExpr{
X: ast.NewIdent(tbl.VarName),
Sel: ast.NewIdent(col.GoFieldName()),
},
},
},
},
}...)
shouldDefineErr = false
}
}
return stmts
}(),
},
}
testList := []ast.Decl{
makeItemFunc,
testCreateUpdateDelete,
@@ -399,13 +561,13 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
Name: ast.NewIdent("."),
},
&ast.ImportSpec{
Path: &ast.BasicLit{Kind: token.STRING, Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"`},
Name: ast.NewIdent("."),
Path: &ast.BasicLit{Kind: token.STRING, Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"`},
},
&ast.ImportSpec{
Path: &ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf(`"%s/pkg/%s"`, gomodName, packageName)},
Name: ast.NewIdent("."),
},
&ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"github.com/go-test/deep"`}},
&ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"github.com/stretchr/testify/assert"`}},
&ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"github.com/stretchr/testify/require"`}},
},

View File

@@ -7,7 +7,7 @@ import (
"os/exec"
"text/template"
. "git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"
)
//go:embed "tpl"
@@ -22,31 +22,31 @@ type PkgOpts struct {
func InitPkg(opts PkgOpts) {
// Run `go mod init`
fmt.Printf("Running... `go mod init %s`\n", opts.ModuleName)
PanicIf(exec.Command("go", "mod", "init", opts.ModuleName).Run())
must.Do(exec.Command("go", "mod", "init", opts.ModuleName).Run())
// Run `git init`, if required
if exec.Command("git", "status").Run() != nil {
// Not in a git repo yet; init one
fmt.Println("Running... `git init`")
PanicIf(exec.Command("git", "init").Run())
must.Do(exec.Command("git", "init").Run())
}
// Create package structure
PanicIf(os.MkdirAll("pkg/db", 0o755))
PanicIf(os.MkdirAll("cmd", 0o755))
PanicIf(os.MkdirAll("doc", 0o755))
PanicIf(os.MkdirAll("sample_data", 0o755))
must.Do(os.MkdirAll("pkg/db", 0o755))
must.Do(os.MkdirAll("cmd", 0o755))
must.Do(os.MkdirAll("doc", 0o755))
must.Do(os.MkdirAll("sample_data", 0o755))
PanicIf(os.WriteFile("pkg/db/schema.sql", Must(tpl.ReadFile("tpl/schema.sql")), 0o664))
PanicIf(os.WriteFile("pkg/db/db.go", Must(tpl.ReadFile("tpl/db.go.tpl")), 0o664))
must.Do(os.WriteFile("pkg/db/schema.sql", must.Get(tpl.ReadFile("tpl/schema.sql")), 0o664))
must.Do(os.WriteFile("pkg/db/db.go", must.Get(tpl.ReadFile("tpl/db.go.tpl")), 0o664))
dbTest := Must(os.Create("pkg/db/db_test.go"))
defer MustClose(dbTest)
t := Must(template.ParseFS(tpl, "tpl/db_test.go.tpl"))
PanicIf(t.Execute(dbTest, opts))
dbTest := must.Get(os.Create("pkg/db/db_test.go"))
defer must.Close(dbTest)
t := must.Get(template.ParseFS(tpl, "tpl/db_test.go.tpl"))
must.Do(t.Execute(dbTest, opts))
PanicIf(os.WriteFile("sample_data/mount.sh", Must(tpl.ReadFile("tpl/mount.sh")), 0o775))
PanicIf(os.WriteFile("sample_data/reset.sh", Must(tpl.ReadFile("tpl/reset.sh")), 0o775))
must.Do(os.WriteFile("sample_data/mount.sh", must.Get(tpl.ReadFile("tpl/mount.sh")), 0o775))
must.Do(os.WriteFile("sample_data/reset.sh", must.Get(tpl.ReadFile("tpl/reset.sh")), 0o775))
// TODO:
// - create `pkg/db/errors.go`

View File

@@ -3,7 +3,7 @@ package db_test
import (
"fmt"
. "git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"
. "{{ .ModuleName }}/pkg/db"
)
@@ -14,6 +14,6 @@ func init() {
TestDB = MakeDB("tmp")
}
func MakeDB(dbName string) *DB {
db := Must(Create(fmt.Sprintf("file:%s?mode=memory&cache=shared", dbName)))
db := must.Get(Create(fmt.Sprintf("file:%s?mode=memory&cache=shared", dbName)))
return db
}

View File

@@ -1,18 +0,0 @@
package flowutils
import "io"
func PanicIf(err error) {
if err != nil {
panic(err)
}
}
func Must[T any](val T, err error) T {
PanicIf(err)
return val
}
func MustClose(closer io.Closer) {
PanicIf(closer.Close())
}

27
pkg/must/must.go Normal file
View File

@@ -0,0 +1,27 @@
package must
import (
"fmt"
"io"
)
func Do(err error) {
if err != nil {
panic(err)
}
}
func Get[T any](val T, err error) T {
Do(err)
return val
}
func Be(condition bool, msg string, args ...any) {
if !condition {
panic(fmt.Errorf(msg, args...)) //nolint:err113 // not a returned error
}
}
func Close(closer io.Closer) {
Do(closer.Close())
}

View File

@@ -9,7 +9,7 @@ import (
"github.com/stretchr/testify/require"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/db"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/schema"
)
@@ -57,7 +57,7 @@ func TestVerifyCorrectMigration(t *testing.T) {
alter table t1 add column data3 integer;
`
db2Config := db.Init(&baseSchema, &[]string{migration})
db2 := flowutils.Must(db2Config.Create(":memory:"))
db2 := must.Get(db2Config.Create(":memory:"))
require.NoError(t, db2Config.CheckAndUpdateVersion(db2))
db2Schema := schema.SchemaFromDB(db2)
@@ -77,7 +77,7 @@ func TestVerifyCorrectMigration(t *testing.T) {
`
db2Config := db.Init(&baseSchema, &[]string{migration1, migration2})
db2 := flowutils.Must(db2Config.Create(":memory:"))
db2 := must.Get(db2Config.Create(":memory:"))
require.NoError(t, db2Config.CheckAndUpdateVersion(db2))
db2Schema := schema.SchemaFromDB(db2)
@@ -93,7 +93,7 @@ func TestIncorrectMigrations(t *testing.T) {
t.Run("missing migration", func(t *testing.T) {
db2Config := db.Init(&baseSchema, &[]string{})
db2 := flowutils.Must(db2Config.Create(":memory:"))
db2 := must.Get(db2Config.Create(":memory:"))
require.NoError(t, db2Config.CheckAndUpdateVersion(db2))
db2Schema := schema.SchemaFromDB(db2)
@@ -114,7 +114,7 @@ func TestIncorrectMigrations(t *testing.T) {
rowid integer primary key
);
`})
db2 := flowutils.Must(db2Config.Create(":memory:"))
db2 := must.Get(db2Config.Create(":memory:"))
require.NoError(t, db2Config.CheckAndUpdateVersion(db2))
db2Schema := schema.SchemaFromDB(db2)
@@ -136,7 +136,7 @@ func TestIncorrectMigrations(t *testing.T) {
);
alter table t1 add column data3 text;
`})
db2 := flowutils.Must(db2Config.Create(":memory:"))
db2 := must.Get(db2Config.Create(":memory:"))
require.NoError(t, db2Config.CheckAndUpdateVersion(db2))
db2Schema := schema.SchemaFromDB(db2)

View File

@@ -10,7 +10,7 @@ import (
"github.com/jmoiron/sqlx"
_ "github.com/mattn/go-sqlite3"
. "git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/textutils"
)
@@ -38,20 +38,20 @@ func SchemaFromDB(db *sqlx.DB) Schema {
db.MustExec(create_views)
var tables []Table
PanicIf(db.Select(&tables, `select name, table_type, is_strict, is_without_rowid from tables`))
must.Do(db.Select(&tables, `select name, table_type, is_strict, is_without_rowid from tables`))
for _, tbl := range tables {
tbl.GoTypeName = TypenameFromTablename(tbl.TableName)
tbl.TypeIDName = tbl.GoTypeName + "ID"
tbl.VarName = strings.ToLower(string(tbl.TableName[0]))
PanicIf(db.Select(&tbl.Columns, `select * from columns where table_name = ?`, tbl.TableName))
must.Do(db.Select(&tbl.Columns, `select * from columns where table_name = ?`, tbl.TableName))
ret.Tables[tbl.TableName] = tbl
}
var indexes []Index
PanicIf(db.Select(&indexes, `select index_name, table_name, is_unique from indexes`))
must.Do(db.Select(&indexes, `select index_name, table_name, is_unique from indexes`))
for _, idx := range indexes {
PanicIf(db.Select(&idx.Columns, `select column_name from index_columns where index_name = ? order by rank`, idx.Name))
must.Do(db.Select(&idx.Columns, `select column_name from index_columns where index_name = ? order by rank`, idx.Name))
ret.Indexes[idx.Name] = idx
}
return ret