29 Commits

Author SHA1 Message Date
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
8c29d455ff codegen: fix defining the 'err' variable multiple times in foreign key checking test
All checks were successful
CI / build-docker (push) Successful in 12s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Successful in 41s
2026-03-19 15:43:31 -07:00
0f9a57dd85 codegen: fix more hardcoded "item" names in generated test file 2026-03-19 15:38:35 -07:00
11fed4b9c7 refactor: create CamelToPascal string helper 2026-03-19 15:35:18 -07:00
4e0836eb2e codegen: fix test generator hardcoding test string in multiple places
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 15s
2026-03-18 13:54:43 -07:00
f82929f6e2 style: remove unused fmtErrorf var
All checks were successful
CI / build-docker (push) Successful in 15s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Successful in 27s
2026-03-18 11:00:18 -07:00
1bc7f9111f codegen: implement "without rowid" tables
Some checks failed
CI / build-docker (push) Successful in 6s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Failing after 16s
2026-03-18 10:39:38 -07:00
c1150954e5 doc: update README
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 19s
2026-03-16 11:45:07 -07:00
fd90830340 doc: add README.md
All checks were successful
CI / build-docker (push) Successful in 11s
CI / build-docker-bootstrap (push) Has been skipped
CI / release-test (push) Successful in 3m33s
2026-03-16 11:38:36 -07:00
21 changed files with 1042 additions and 580 deletions

View File

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

View File

@@ -60,8 +60,6 @@ linters:
- -ST1000 # Re-enable this once we have docstrings - -ST1000 # Re-enable this once we have docstrings
- -ST1003 # I like snake_case - -ST1003 # I like snake_case
- -ST1013 # HTTP status codes are shorter and more readable than names - -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: exclusions:
generated: lax # Don't lint generated files generated: lax # Don't lint generated files
paths: paths:

48
README.md Normal file
View File

@@ -0,0 +1,48 @@
# GAS stack
## Compiling
Requires a Go compiler (minimum 1.22.5) and a C compiler, due to use of CGo.
```sh
git clone https://git.offline-twitter.com/offline-labs/gas-stack.git
cd gas-stack
go build -o gas -tags fts5 ./cmd
# Installation (optional)
sudo mv gas /usr/local/bin # ...or anywhere on your $PATH
which gas # should print "/usr/local/bin/gas"
```
## Using
The linter (`gas sqlite_lint`) is stable and useful.
The code generator is buggy, incomplete, and not remotely stable, but still quite useful. Don't expect it to produce perfectly working code, or even to compile correctly (e.g., you'll probably have to fix the imports). Copy-paste the parts that are useful, and delete the parts that aren't.
#### Linter
```sh
gas sqlite_lint <path/to/schema.sql>
```
#### Code generator
```sh
gas generate table_name # Generates a model
gas generate --test table_name # Optional: generates tests
```
It prints to the console. You can copy-paste the result. Or you can use bash redirection:
```sh
gas generate users > pkg/db/user.go
gas generate --test users > pkg/db/user_test.go
```
Useful flags:
- `--schema`: by default, `gas generate` assumes that the schema is at `pkg/db/schema.sql`. Use `gas generate --schema <path/to/schema.sql> [...]` to indicate otherwise

View File

@@ -6,11 +6,12 @@ import (
"go/ast" "go/ast"
"go/token" "go/token"
"os" "os"
"slices"
"github.com/spf13/cobra" "github.com/spf13/cobra"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/codegen/modelgenerate" "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" "git.offline-twitter.com/offline-labs/gas-stack/pkg/schema"
) )
@@ -23,22 +24,22 @@ var generate_model = &cobra.Command{
Args: cobra.ExactArgs(1), Args: cobra.ExactArgs(1),
RunE: func(cmd *cobra.Command, args []string) error { RunE: func(cmd *cobra.Command, args []string) error {
path := Must(cmd.Flags().GetString("schema")) path := must.Get(cmd.Flags().GetString("schema"))
modname := Must(cmd.Flags().GetString("modname")) modname := must.Get(cmd.Flags().GetString("modname"))
sql, err := os.ReadFile(path) sql, err := os.ReadFile(path)
if err != nil { if err != nil {
return fmt.Errorf("reading path %s: %w", path, err) return fmt.Errorf("reading path %s: %w", path, err)
} }
db := schema.InitDB(string(sql)) db := schema.InitDB(string(sql))
schema := schema.SchemaFromDB(db) sch := schema.SchemaFromDB(db)
table, isOk := schema.Tables[args[0]] table, isOk := sch.Tables[args[0]]
if !isOk { if !isOk {
return ErrNoSuchTable return ErrNoSuchTable
} }
if Must(cmd.Flags().GetBool("test")) { if must.Get(cmd.Flags().GetBool("test")) {
file2 := modelgenerate.GenerateModelTestAST(table, schema, modname) file2 := modelgenerate.GenerateModelTestAST(table, sch, modname)
PanicIf(modelgenerate.FprintWithComments(os.Stdout, file2)) must.Do(modelgenerate.FprintWithComments(os.Stdout, file2))
} else { } else {
decls := []ast.Decl{ decls := []ast.Decl{
&ast.GenDecl{ &ast.GenDecl{
@@ -46,16 +47,14 @@ var generate_model = &cobra.Command{
Specs: []ast.Spec{ Specs: []ast.Spec{
&ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"database/sql"`}}, &ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"database/sql"`}},
&ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"errors"`}}, &ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"errors"`}},
&ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"fmt"`}},
&ast.ImportSpec{ &ast.ImportSpec{
Name: ast.NewIdent("."), Name: ast.NewIdent("."),
Path: &ast.BasicLit{Kind: token.STRING, Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/db"`}, Path: &ast.BasicLit{Kind: token.STRING, Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/db"`},
}, },
&ast.ImportSpec{ &ast.ImportSpec{
Name: ast.NewIdent("."),
Path: &ast.BasicLit{ Path: &ast.BasicLit{
Kind: token.STRING, 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"`,
}, },
}, },
}, },
@@ -69,15 +68,33 @@ var generate_model = &cobra.Command{
modelgenerate.GenerateSQLFieldsConst(table), modelgenerate.GenerateSQLFieldsConst(table),
modelgenerate.GenerateSaveItemFunc(table), modelgenerate.GenerateSaveItemFunc(table),
modelgenerate.GenerateDeleteItemFunc(table), modelgenerate.GenerateDeleteItemFunc(table),
)
if table.IsWithoutRowid {
decls = append(decls,
modelgenerate.GenerateGetItemBy(table, table.PrimaryKeyColumns()),
)
} else {
decls = append(decls,
modelgenerate.GenerateGetItemByIDFunc(table), modelgenerate.GenerateGetItemByIDFunc(table),
) )
for _, index := range schema.Indexes { }
for _, index := range sch.Indexes {
if index.TableName != table.TableName { if index.TableName != table.TableName {
// Skip indexes on other tables // Skip indexes on other tables
continue continue
} }
if index.IsUnique && len(index.Columns) == 1 { if slices.Contains(index.Columns, "") {
decls = append(decls, modelgenerate.GenerateGetItemByUniqColFunc(table, table.GetColumnByName(index.Columns[0]))) // 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)) decls = append(decls, modelgenerate.GenerateGetAllItemsFunc(table))
@@ -87,7 +104,7 @@ var generate_model = &cobra.Command{
Decls: decls, Decls: decls,
} }
PanicIf(modelgenerate.FprintWithComments(os.Stdout, file)) must.Do(modelgenerate.FprintWithComments(os.Stdout, file))
} }
return nil return nil

View File

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

View File

@@ -46,11 +46,19 @@ create table items (
created_at integer not null, created_at integer not null,
updated_at integer not null updated_at integer not null
) strict; ) strict;
create table item_to_item (
item1_id integer references items(rowid),
item2_id integer references items(rowid),
primary key (item1_id, item2_id)
) strict, without rowid;
EOF EOF
# Generate an item model and test file # Generate an item model and test file
$gas generate items > pkg/db/item.go $gas generate items > pkg/db/item.go
$gas generate items --test > pkg/db/item_test.go $gas generate items --test > pkg/db/item_test.go
$gas generate item_to_item > pkg/db/item_to_item.go
go mod tidy go mod tidy
# Run the tests # Run the tests

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. // we don't have to do that.
var TrailingComments = map[ast.Node]string{} 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 { func mustCall(inner ast.Expr) *ast.CallExpr {
return &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}, Args: []ast.Expr{inner},
} }
} }
@@ -73,7 +81,7 @@ func FprintWithComments(w io.Writer, file *ast.File) error {
fset = token.NewFileSet() fset = token.NewFileSet()
parsed, err := parser.ParseFile(fset, "", buf.Bytes(), parser.ParseComments) parsed, err := parser.ParseFile(fset, "", buf.Bytes(), parser.ParseComments)
if err != nil { 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. // Parallel walk: apply TrailingComments from the side map.

View File

@@ -5,6 +5,7 @@ import (
"fmt" "fmt"
"go/ast" "go/ast"
"go/token" "go/token"
"slices"
"strings" "strings"
"github.com/jinzhu/inflection" "github.com/jinzhu/inflection"
@@ -20,11 +21,10 @@ import (
var ( var (
dbRecv = &ast.FieldList{List: []*ast.Field{{Names: []*ast.Ident{ast.NewIdent("db")}, Type: ast.NewIdent("DB")}}} dbRecv = &ast.FieldList{List: []*ast.Field{{Names: []*ast.Ident{ast.NewIdent("db")}, Type: ast.NewIdent("DB")}}}
dbDB = &ast.SelectorExpr{X: ast.NewIdent("db"), Sel: ast.NewIdent("DB")} dbDB = &ast.SelectorExpr{X: ast.NewIdent("db"), Sel: ast.NewIdent("DB")}
fmtErrorf = &ast.SelectorExpr{X: ast.NewIdent("fmt"), Sel: ast.NewIdent("Errorf")}
) )
func SQLFieldsConstIdent(tbl schema.Table) *ast.Ident { 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. // GoTypeForColumn returns a type expression for this column.
@@ -54,6 +54,50 @@ func GoTypeForColumn(c schema.Column) ast.Expr {
} }
} }
// 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.EQL,
Y: &ast.BasicLit{Kind: token.INT, Value: "1"},
},
&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,
}}
}
// --------------- // ---------------
// Generators // Generators
// --------------- // ---------------
@@ -146,9 +190,21 @@ func buildFKCheckLambda(tbl schema.Table) (*ast.AssignStmt, bool) {
structFieldName := col.GoFieldName() structFieldName := col.GoFieldName()
structField := &ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent(structFieldName)} 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() { if col.IsNonCodeTableForeignKey() {
// Real foreign key; look up referent by ID to see if it exists // 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{ Init: &ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("_"), ast.NewIdent("err")}, Lhs: []ast.Expr{ast.NewIdent("_"), ast.NewIdent("err")},
Tok: token.DEFINE, Tok: token.DEFINE,
@@ -182,10 +238,10 @@ func buildFKCheckLambda(tbl schema.Table) (*ast.AssignStmt, bool) {
}) })
} else { } else {
// Code table value. Query the table to see if it exists // Code table value. Query the table to see if it exists
ret = append(ret, &ast.IfStmt{ return wrap(&ast.IfStmt{
Init: &ast.AssignStmt{ Init: &ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("err")}, Lhs: []ast.Expr{ast.NewIdent("err")},
Tok: token.ASSIGN, Tok: token.DEFINE,
Rhs: []ast.Expr{ Rhs: []ast.Expr{
&ast.CallExpr{ &ast.CallExpr{
Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("Get")}, Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("Get")},
@@ -219,6 +275,7 @@ func buildFKCheckLambda(tbl schema.Table) (*ast.AssignStmt, bool) {
}, },
}) })
} }
}())
} }
// final return nil // final return nil
ret = append(ret, &ast.ReturnStmt{Results: []ast.Expr{ast.NewIdent("nil")}}) ret = append(ret, &ast.ReturnStmt{Results: []ast.Expr{ast.NewIdent("nil")}})
@@ -255,10 +312,29 @@ func GenerateSaveItemFunc(tbl schema.Table) *ast.FuncDecl {
if col.Name == "created_at" && hasCreatedAt { if col.Name == "created_at" && hasCreatedAt {
continue continue
} }
if !col.IsPrimaryKey { // Don't try to update primary key columns (mainly for w/o rowid tables)
updatePairs = append(updatePairs, col.Name+"="+val) updatePairs = append(updatePairs, col.Name+"="+val)
} }
insertStmt := fmt.Sprintf("\n\t\t insert into %s (%s)\n\t\t values (%s)\n\t\t", tbl.TableName, strings.Join(insertCols, ", "), strings.Join(insertVals, ", ")) }
updateStmt := fmt.Sprintf("\n\t\t update %s\n\t\t set %s\n\t\t where rowid = :rowid\n\t\t", tbl.TableName, strings.Join(updatePairs, ",\n\t\t ")) insertStmt := fmt.Sprintf("\n\t\t insert into %s (%s)\n\t\t values (%s)\n\t\t",
tbl.TableName,
strings.Join(insertCols, ", "),
strings.Join(insertVals, ", "),
)
updateStmt := fmt.Sprintf("\n\t\t update %s\n\t\t set %s\n\t\t where rowid = :rowid\n\t\t",
tbl.TableName,
strings.Join(updatePairs, ",\n\t\t "),
)
upsertStmt := fmt.Sprintf("\n\t insert into %s (%s)\n\t values (%s)\n\t",
tbl.TableName,
strings.Join(insertCols, ", "),
strings.Join(insertVals, ", "),
)
if len(updatePairs) == 0 {
upsertStmt = upsertStmt + " on conflict do nothing\n\t"
} else {
upsertStmt = upsertStmt + fmt.Sprintf(" on conflict do update\n\t set %s\n\t", strings.Join(updatePairs, ",\n\t "))
}
checkForeignKeyFailuresAssignment, hasFks := buildFKCheckLambda(tbl) checkForeignKeyFailuresAssignment, hasFks := buildFKCheckLambda(tbl)
@@ -276,43 +352,25 @@ func GenerateSaveItemFunc(tbl schema.Table) *ast.FuncDecl {
Rhs: []ast.Expr{&ast.CallExpr{Fun: ast.NewIdent("TimestampNow"), Args: []ast.Expr{}}}, Rhs: []ast.Expr{&ast.CallExpr{Fun: ast.NewIdent("TimestampNow"), Args: []ast.Expr{}}},
}) })
} }
// if item.ID == 0 {...} else {...}
ret = append(ret, &ast.IfStmt{ namedExecStmt := func(stmt string) []ast.Stmt {
Cond: &ast.BinaryExpr{ queryStmt := &ast.CallExpr{
X: &ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("ID")},
Op: token.EQL,
Y: &ast.BasicLit{Kind: token.INT, Value: "0"},
},
Body: &ast.BlockStmt{
// Do create
List: append(
func() []ast.Stmt {
ret1 := []ast.Stmt{Comment("Do create")}
if hasCreatedAt {
// Auto-timestamps: created_at
ret1 = append(ret1, &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{}}},
})
}
namedExecStmt := &ast.CallExpr{
Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("NamedExec")}, Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("NamedExec")},
Args: []ast.Expr{ Args: []ast.Expr{
&ast.BasicLit{Kind: token.STRING, Value: "`" + insertStmt + "`"}, &ast.BasicLit{Kind: token.STRING, Value: "`" + stmt + "`"},
ast.NewIdent(tbl.VarName), ast.NewIdent(tbl.VarName),
}, },
} }
if !hasFks { if !hasFks {
// No foreign key checking needed; just use `Must` for brevity // No foreign key checking needed; just use `must.Get` for brevity
return append(ret1, &ast.AssignStmt{ return []ast.Stmt{&ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("result")}, Lhs: []ast.Expr{ast.NewIdent("result")},
Tok: token.DEFINE, Tok: token.DEFINE,
Rhs: []ast.Expr{mustCall(namedExecStmt)}, Rhs: []ast.Expr{mustCall(queryStmt)},
}) }}
} }
// There's foreign keys
return append(ret1, return []ast.Stmt{
// result, err := db.DB.NamedExec(`...`, u) // result, err := db.DB.NamedExec(`...`, u)
&ast.AssignStmt{ &ast.AssignStmt{
Lhs: []ast.Expr{ Lhs: []ast.Expr{
@@ -320,9 +378,8 @@ func GenerateSaveItemFunc(tbl schema.Table) *ast.FuncDecl {
ast.NewIdent("err"), ast.NewIdent("err"),
}, },
Tok: token.DEFINE, Tok: token.DEFINE,
Rhs: []ast.Expr{namedExecStmt}, Rhs: []ast.Expr{queryStmt},
}, },
// if fkErr := checkForeignKeyFailures(err); fkErr != nil { return fkErr } else if err != nil { panic(err) } // if fkErr := checkForeignKeyFailures(err); fkErr != nil { return fkErr } else if err != nil { panic(err) }
&ast.IfStmt{ &ast.IfStmt{
Init: &ast.AssignStmt{ Init: &ast.AssignStmt{
@@ -367,7 +424,63 @@ 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, MustBeRowsAffected(tbl))
} else {
// if item.ID == 0 {...} else {...}
ret = append(ret, &ast.IfStmt{
Cond: &ast.BinaryExpr{
X: &ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("ID")},
Op: token.EQL,
Y: &ast.BasicLit{Kind: token.INT, Value: "0"},
},
Body: &ast.BlockStmt{
// Do create
List: append(
func() []ast.Stmt {
ret1 := []ast.Stmt{Comment("Do create")}
if hasCreatedAt {
// Auto-timestamps: created_at
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)...)
}(), }(),
&ast.AssignStmt{ &ast.AssignStmt{
Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("ID")}}, Lhs: []ast.Expr{&ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("ID")}},
@@ -384,31 +497,16 @@ func GenerateSaveItemFunc(tbl schema.Table) *ast.FuncDecl {
}, },
Else: &ast.BlockStmt{ Else: &ast.BlockStmt{
// Do update // Do update
List: []ast.Stmt{ List: append(
Comment("Do update"), []ast.Stmt{Comment("Do update")},
&ast.AssignStmt{ append(
Lhs: []ast.Expr{ast.NewIdent("result")}, namedExecStmt(updateStmt),
Tok: token.DEFINE, MustBeRowsAffected(tbl),
Rhs: []ast.Expr{mustCall(&ast.CallExpr{ )...,
Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("NamedExec")}, ),
Args: []ast.Expr{&ast.BasicLit{Kind: token.STRING, Value: "`" + updateStmt + "`"}, ast.NewIdent(tbl.VarName)},
})},
},
&ast.IfStmt{
Cond: &ast.BinaryExpr{
X: mustCall(&ast.CallExpr{
Fun: &ast.SelectorExpr{X: ast.NewIdent("result"), Sel: ast.NewIdent("RowsAffected")},
Args: []ast.Expr{},
}),
Op: token.NEQ,
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.CallExpr{Fun: fmtErrorf, Args: []ast.Expr{&ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf("\"got %s with ID (%%d), so attempted update, but it doesn't exist\"", strings.ToLower(tbl.GoTypeName))}, &ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("ID")}}}}}}}},
},
},
}, },
}) })
}
if hasFks { if hasFks {
// If there's foreign key checking, it needs to return an error (or nil) // If there's foreign key checking, it needs to return an error (or nil)
ret = append(ret, &ast.ReturnStmt{Results: []ast.Expr{ast.NewIdent("nil")}}) ret = append(ret, &ast.ReturnStmt{Results: []ast.Expr{ast.NewIdent("nil")}})
@@ -442,6 +540,97 @@ func getByIDFuncName(tblname string) string {
return "Get" + schema.TypenameFromTablename(tblname) + "ByID" 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 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 `"},
Op: token.ADD,
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 "))},
}
return &ast.FuncDecl{
Recv: dbRecv,
Name: ast.NewIdent(fmt.Sprintf("Get%sBy%s", schema.TypenameFromTablename(tbl.TableName), strings.Join(funcNameSuffix, "And"))),
Type: &ast.FuncType{
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")},
}},
},
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: append([]ast.Expr{&ast.UnaryExpr{Op: token.AND, X: ast.NewIdent("ret")}, selectExpr}, sqlParams...),
}},
},
&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.ReturnStmt{},
},
},
}
}
// GenerateGetItemByIDFunc produces an AST for the `GetXyzByID()` function. // GenerateGetItemByIDFunc produces an AST for the `GetXyzByID()` function.
// E.g., a table with `table.TypeName = "foods"` will produce a "GetFoodByID()" function. // E.g., a table with `table.TypeName = "foods"` will produce a "GetFoodByID()" function.
func GenerateGetItemByIDFunc(tbl schema.Table) *ast.FuncDecl { func GenerateGetItemByIDFunc(tbl schema.Table) *ast.FuncDecl {
@@ -474,19 +663,41 @@ func GenerateGetItemByIDFunc(tbl schema.Table) *ast.FuncDecl {
}, },
} }
funcDecl := &ast.FuncDecl{ return &ast.FuncDecl{
Recv: dbRecv, Recv: dbRecv,
Name: ast.NewIdent(getByIDFuncName(tbl.TableName)), Name: ast.NewIdent(getByIDFuncName(tbl.TableName)),
Type: &ast.FuncType{Params: arg, Results: result}, Type: &ast.FuncType{Params: arg, Results: result},
Body: funcBody, Body: funcBody,
} }
return funcDecl
} }
// GenerateGetItemByUniqColFunc produces an AST for the `GetXyzByID()` function. // GenerateGetItemByUniqColFunc produces an AST for the `GetXyzByCol()` function, for a unique
// E.g., a table with `table.TypeName = "foods"` will produce a "GetFoodByID()" function. // 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 { 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{ selectExpr := &ast.BinaryExpr{
X: &ast.BinaryExpr{ X: &ast.BinaryExpr{
X: &ast.BasicLit{Kind: token.STRING, Value: "`\n\t select `"}, X: &ast.BasicLit{Kind: token.STRING, Value: "`\n\t select `"},
@@ -494,40 +705,40 @@ func GenerateGetItemByUniqColFunc(tbl schema.Table, col schema.Column) *ast.Func
Y: SQLFieldsConstIdent(tbl), Y: SQLFieldsConstIdent(tbl),
}, },
Op: token.ADD, 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{ return &ast.FuncDecl{
Recv: dbRecv, 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{ Type: &ast.FuncType{
Params: &ast.FieldList{List: []*ast.Field{ Params: funcParams,
{Names: []*ast.Ident{param}, Type: GoTypeForColumn(col)},
}},
Results: &ast.FieldList{List: []*ast.Field{ Results: &ast.FieldList{List: []*ast.Field{
{Names: []*ast.Ident{ast.NewIdent("ret")}, Type: ast.NewIdent(tbl.GoTypeName)}, {Names: []*ast.Ident{ast.NewIdent("ret")}, Type: &ast.ArrayType{Elt: ast.NewIdent(tbl.GoTypeName)}},
{Names: []*ast.Ident{ast.NewIdent("err")}, Type: ast.NewIdent("error")},
}}, }},
}, },
Body: &ast.BlockStmt{ Body: &ast.BlockStmt{
List: []ast.Stmt{ List: []ast.Stmt{
&ast.AssignStmt{ &ast.ExprStmt{X: selectCall},
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.ReturnStmt{}, &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. // GenerateGetAllItemsFunc produces an AST for the `GetAllXyzs()` function.
// E.g., a table with `table.TypeName = "foods"` will produce a "GetAllFoods()" function. // E.g., a table with `table.TypeName = "foods"` will produce a "GetAllFoods()" function.
func GenerateGetAllItemsFunc(tbl schema.Table) *ast.FuncDecl { func GenerateGetAllItemsFunc(tbl schema.Table) *ast.FuncDecl {
@@ -536,10 +747,7 @@ func GenerateGetAllItemsFunc(tbl schema.Table) *ast.FuncDecl {
{Names: []*ast.Ident{ast.NewIdent("ret")}, Type: &ast.ArrayType{Elt: ast.NewIdent(tbl.GoTypeName)}}, {Names: []*ast.Ident{ast.NewIdent("ret")}, Type: &ast.ArrayType{Elt: ast.NewIdent(tbl.GoTypeName)}},
}} }}
selectCall := &ast.CallExpr{ selectCall := doCall(&ast.CallExpr{
Fun: ast.NewIdent("PanicIf"),
Args: []ast.Expr{
&ast.CallExpr{
Fun: &ast.SelectorExpr{ Fun: &ast.SelectorExpr{
X: dbDB, X: dbDB,
Sel: ast.NewIdent("Select"), Sel: ast.NewIdent("Select"),
@@ -556,9 +764,7 @@ func GenerateGetAllItemsFunc(tbl schema.Table) *ast.FuncDecl {
Y: &ast.BasicLit{Kind: token.STRING, Value: "` from " + tbl.TableName + "`"}, Y: &ast.BasicLit{Kind: token.STRING, Value: "` from " + tbl.TableName + "`"},
}, },
}, },
}, })
},
}
funcBody := &ast.BlockStmt{ funcBody := &ast.BlockStmt{
List: []ast.Stmt{ List: []ast.Stmt{
@@ -581,10 +787,12 @@ func GenerateGetAllItemsFunc(tbl schema.Table) *ast.FuncDecl {
// GenerateDeleteItemFunc produces an AST for the `DeleteXyz()` function. // GenerateDeleteItemFunc produces an AST for the `DeleteXyz()` function.
// E.g., a table with `table.TypeName = "foods"` will produce a "DeleteFood()" function. // E.g., a table with `table.TypeName = "foods"` will produce a "DeleteFood()" function.
func GenerateDeleteItemFunc(tbl schema.Table) *ast.FuncDecl { func GenerateDeleteItemFunc(tbl schema.Table) *ast.FuncDecl {
arg := &ast.FieldList{List: []*ast.Field{{ colNames := []string{}
Names: []*ast.Ident{ast.NewIdent(tbl.VarName)}, for _, c := range tbl.PrimaryKeyColumns() {
Type: ast.NewIdent(tbl.GoTypeName), colNames = append(colNames, fmt.Sprintf("%s = :%s", c.Name, c.Name))
}}} }
sqlStr := "`delete from " + tbl.TableName + fmt.Sprintf(" where %s`", strings.Join(colNames, " and "))
funcBody := &ast.BlockStmt{ funcBody := &ast.BlockStmt{
List: []ast.Stmt{ List: []ast.Stmt{
@@ -592,41 +800,24 @@ func GenerateDeleteItemFunc(tbl schema.Table) *ast.FuncDecl {
Lhs: []ast.Expr{ast.NewIdent("result")}, Lhs: []ast.Expr{ast.NewIdent("result")},
Tok: token.DEFINE, Tok: token.DEFINE,
Rhs: []ast.Expr{mustCall(&ast.CallExpr{ Rhs: []ast.Expr{mustCall(&ast.CallExpr{
Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("Exec")}, Fun: &ast.SelectorExpr{X: dbDB, Sel: ast.NewIdent("NamedExec")},
Args: []ast.Expr{ Args: []ast.Expr{
&ast.BasicLit{Kind: token.STRING, Value: "`delete from " + tbl.TableName + " where rowid = ?`"}, &ast.BasicLit{Kind: token.STRING, Value: sqlStr},
&ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("ID")}, ast.NewIdent(tbl.VarName),
}, },
})}, })},
}, },
&ast.IfStmt{ MustBeRowsAffected(tbl),
Cond: &ast.BinaryExpr{
X: mustCall(
&ast.CallExpr{Fun: &ast.SelectorExpr{X: ast.NewIdent("result"), Sel: ast.NewIdent("RowsAffected")}, Args: []ast.Expr{}},
),
Op: token.NEQ,
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.CallExpr{
Fun: fmtErrorf,
Args: []ast.Expr{
&ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf("\"tried to delete %s with ID (%%d) but it doesn't exist\"", strings.ToLower(tbl.GoTypeName))},
&ast.SelectorExpr{X: ast.NewIdent(tbl.VarName), Sel: ast.NewIdent("ID")},
},
}},
}},
}},
},
}, },
} }
funcDecl := &ast.FuncDecl{ funcDecl := &ast.FuncDecl{
Recv: dbRecv, Recv: dbRecv,
Name: ast.NewIdent("Delete" + tbl.GoTypeName), Name: ast.NewIdent("Delete" + tbl.GoTypeName),
Type: &ast.FuncType{Params: arg, Results: nil}, Type: &ast.FuncType{Params: &ast.FieldList{List: []*ast.Field{{
Names: []*ast.Ident{ast.NewIdent(tbl.VarName)},
Type: ast.NewIdent(tbl.GoTypeName),
}}}, Results: nil},
Body: funcBody, Body: funcBody,
} }
return funcDecl return funcDecl

View File

@@ -4,20 +4,102 @@ import (
"fmt" "fmt"
"go/ast" "go/ast"
"go/token" "go/token"
"slices"
"strings"
"github.com/jinzhu/inflection" "github.com/jinzhu/inflection"
pkgschema "git.offline-twitter.com/offline-labs/gas-stack/pkg/schema" pkgschema "git.offline-twitter.com/offline-labs/gas-stack/pkg/schema"
"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. // 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 { func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodName string) *ast.File {
packageName := "db" packageName := "db"
testpackageName := packageName + "_test" testpackageName := packageName + "_test"
makeHelperName := ast.NewIdent("Make" + tbl.GoTypeName)
hasCreatedAt, hasUpdatedAt := tbl.HasAutoTimestamps()
updateTestFields := UpdateTestFields(tbl)
// func MakeItem() Item { return Item{} } // func MakeItem() Item { return Item{} }
makeItemFunc := &ast.FuncDecl{ makeItemFunc := &ast.FuncDecl{
Name: ast.NewIdent("Make" + tbl.GoTypeName), Name: makeHelperName,
Type: &ast.FuncType{ Type: &ast.FuncType{
Params: &ast.FieldList{}, Params: &ast.FieldList{},
Results: &ast.FieldList{ Results: &ast.FieldList{
@@ -32,24 +114,15 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
Results: []ast.Expr{ Results: []ast.Expr{
&ast.CompositeLit{ &ast.CompositeLit{
Type: ast.NewIdent(tbl.GoTypeName), Type: ast.NewIdent(tbl.GoTypeName),
Elts: []ast.Expr{ Elts: func() (ret []ast.Expr) {
&ast.KeyValueExpr{ for _, c := range updateTestFields {
Key: ast.NewIdent("Data"), ret = append(ret, &ast.KeyValueExpr{
Value: &ast.CompositeLit{ Key: ast.NewIdent(c.GoFieldName()),
Type: &ast.ArrayType{ Value: SampleValue(c, 0),
Elt: ast.NewIdent("byte"), })
}, }
Elts: []ast.Expr{}, return
}, }(),
},
&ast.KeyValueExpr{
Key: ast.NewIdent("Description"),
Value: &ast.BasicLit{
Kind: token.STRING,
Value: `""`,
},
},
},
}, },
}, },
}, },
@@ -57,14 +130,74 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
}, },
} }
testObj := ast.NewIdent("item") testObj := ast.NewIdent(textutils.CamelToPascal(tbl.GoTypeName))
testObj2 := ast.NewIdent("item2") testObj2 := ast.NewIdent(textutils.CamelToPascal(tbl.GoTypeName) + "2")
fieldName := ast.NewIdent("Description") // TODO
description1 := `"an item"`
description2 := `"a big item"`
testDB := ast.NewIdent("TestDB") 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{ testFuncType := &ast.FuncType{
Params: &ast.FieldList{ Params: &ast.FieldList{
@@ -75,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{ testCreateUpdateDelete := &ast.FuncDecl{
Name: ast.NewIdent("TestCreateUpdateDelete" + tbl.GoTypeName), Name: ast.NewIdent("TestCreateUpdateDelete" + tbl.GoTypeName),
Type: testFuncType, Type: testFuncType,
@@ -94,36 +324,47 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
&ast.AssignStmt{ &ast.AssignStmt{
Lhs: []ast.Expr{testObj}, Lhs: []ast.Expr{testObj},
Tok: token.DEFINE, Tok: token.DEFINE,
Rhs: []ast.Expr{&ast.CallExpr{Fun: ast.NewIdent("MakeItem"), Args: nil}}, Rhs: []ast.Expr{&ast.CallExpr{Fun: makeHelperName, Args: nil}},
}, },
// item.Description = "an item" // item.Description = <sample value, offset 1>
&ast.AssignStmt{ // ...one assignment per updatable field
}
for _, c := range updateTestFields {
stmts = append(stmts, &ast.AssignStmt{
Lhs: []ast.Expr{ Lhs: []ast.Expr{
&ast.SelectorExpr{ &ast.SelectorExpr{
X: ast.NewIdent("item"), X: testObj,
Sel: ast.NewIdent("Description"), Sel: ast.NewIdent(c.GoFieldName()),
}, },
}, },
Tok: token.ASSIGN, Tok: token.ASSIGN,
Rhs: []ast.Expr{ Rhs: []ast.Expr{SampleValue(c, 1)},
&ast.BasicLit{ })
Kind: token.STRING, }
Value: `"an item"`, stmts = append(stmts,
}, // TestDB.SaveItem(&item), possibly with error check
}, &ast.ExprStmt{X: func() *ast.CallExpr {
}, mainExpr := &ast.CallExpr{
// TestDB.SaveItem(&item)
&ast.ExprStmt{X: &ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Save" + tbl.GoTypeName)}, Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Save" + tbl.GoTypeName)},
Args: []ast.Expr{&ast.UnaryExpr{Op: token.AND, X: testObj}}, 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) // 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")}, 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")}}, Args: []ast.Expr{ast.NewIdent("t"), &ast.SelectorExpr{X: testObj, Sel: ast.NewIdent("ID")}},
}}, }})
} }
// After create: assert timestamps are set // After create: assert timestamps are set
@@ -138,63 +379,48 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
BlankLine(), BlankLine(),
Comment("Load"), Comment("Load"),
// item2 := Must(TestDB.GetItemByID(item.ID)) // item2 := must.Get(TestDB.GetItemByID(item.ID))
&ast.AssignStmt{ &ast.AssignStmt{
Lhs: []ast.Expr{testObj2}, Lhs: []ast.Expr{testObj2},
Tok: token.DEFINE, Tok: token.DEFINE,
Rhs: []ast.Expr{mustCall(&ast.CallExpr{ Rhs: []ast.Expr{mustCall(getItemByPKCall(testObj))},
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Get" + tbl.GoTypeName + "ByID")},
Args: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: ast.NewIdent("ID")}},
})},
}, },
// assert.Equal(t, "an item", item2.Description) // if deep.Equal(...) {...}
&ast.ExprStmt{X: &ast.CallExpr{ makeDeepEqual(testObj, testObj2),
Fun: &ast.SelectorExpr{X: ast.NewIdent("assert"), Sel: ast.NewIdent("Equal")},
Args: []ast.Expr{
ast.NewIdent("t"),
&ast.BasicLit{Kind: token.STRING, Value: description1},
&ast.SelectorExpr{X: testObj2, Sel: fieldName},
},
}},
) )
stmts = append(stmts, stmts = append(stmts,
BlankLine(), BlankLine(),
Comment("Update"), Comment("Update"),
)
// item.Description = "a big item" // item.Description = <sample value, offset 2>
&ast.AssignStmt{ // ...one assignment per updatable field
Lhs: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: fieldName}}, for _, c := range updateTestFields {
stmts = append(stmts, &ast.AssignStmt{
Lhs: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: ast.NewIdent(c.GoFieldName())}},
Tok: token.ASSIGN, 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) // TestDB.SaveItem(&item)
&ast.ExprStmt{X: &ast.CallExpr{ &ast.ExprStmt{X: &ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Save" + tbl.GoTypeName)}, Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Save" + tbl.GoTypeName)},
Args: []ast.Expr{&ast.UnaryExpr{Op: token.AND, X: testObj}}, 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{ &ast.AssignStmt{
Lhs: []ast.Expr{testObj2}, Lhs: []ast.Expr{testObj2},
Tok: token.ASSIGN, Tok: token.ASSIGN,
Rhs: []ast.Expr{mustCall(&ast.CallExpr{ Rhs: []ast.Expr{mustCall(getItemByPKCall(testObj))},
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Get" + tbl.GoTypeName + "ByID")},
Args: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: ast.NewIdent("ID")}},
})},
}, },
// assert.Equal(t, item.Description, item2.Description) // if deep.Equal(...) {...}
&ast.ExprStmt{X: &ast.CallExpr{ makeDeepEqual(testObj, testObj2),
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},
},
}},
) )
indexGets, hasIndexedGets := []ast.Stmt{ indexGets, hasIndexedGets := []ast.Stmt{
@@ -207,8 +433,21 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
// Skip indexes on other tables // Skip indexes on other tables
continue continue
} }
if index.IsUnique && len(index.Columns) == 1 { if slices.Contains(index.Columns, "") {
col := tbl.GetColumnByName(index.Columns[0]) // 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{ indexGets = append(indexGets, []ast.Stmt{
// assert.Equal(t, item2, TestDB.GetItemByXYZ(...)) // assert.Equal(t, item2, TestDB.GetItemByXYZ(...))
&ast.ExprStmt{X: &ast.CallExpr{ // TODO: what if just delete the "ExprStmt" wrapper? &ast.ExprStmt{X: &ast.CallExpr{ // TODO: what if just delete the "ExprStmt" wrapper?
@@ -218,16 +457,32 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
testObj2, testObj2,
mustCall(&ast.CallExpr{ mustCall(&ast.CallExpr{
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent( 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]))) } else {
hasIndexedGets = true 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 { if hasIndexedGets {
@@ -248,10 +503,7 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
&ast.AssignStmt{ &ast.AssignStmt{
Lhs: []ast.Expr{ast.NewIdent("_"), ast.NewIdent("err")}, Lhs: []ast.Expr{ast.NewIdent("_"), ast.NewIdent("err")},
Tok: token.DEFINE, Tok: token.DEFINE,
Rhs: []ast.Expr{&ast.CallExpr{ Rhs: []ast.Expr{getItemByPKCall(testObj)},
Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent("Get" + tbl.GoTypeName + "ByID")},
Args: []ast.Expr{&ast.SelectorExpr{X: testObj, Sel: ast.NewIdent("ID")}},
}},
}, },
// assert.ErrorIs(t, err, db.ErrNotInDB) // assert.ErrorIs(t, err, db.ErrNotInDB)
@@ -289,92 +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),
},
},
},
}
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: token.DEFINE,
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()),
},
},
},
},
}...)
}
}
return stmts
}(),
},
}
testList := []ast.Decl{ testList := []ast.Decl{
makeItemFunc, makeItemFunc,
testCreateUpdateDelete, testCreateUpdateDelete,
@@ -395,13 +561,13 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam
Name: ast.NewIdent("."), Name: ast.NewIdent("."),
}, },
&ast.ImportSpec{ &ast.ImportSpec{
Path: &ast.BasicLit{Kind: token.STRING, Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils"`}, Path: &ast.BasicLit{Kind: token.STRING, Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"`},
Name: ast.NewIdent("."),
}, },
&ast.ImportSpec{ &ast.ImportSpec{
Path: &ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf(`"%s/pkg/%s"`, gomodName, packageName)}, Path: &ast.BasicLit{Kind: token.STRING, Value: fmt.Sprintf(`"%s/pkg/%s"`, gomodName, packageName)},
Name: ast.NewIdent("."), 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/assert"`}},
&ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"github.com/stretchr/testify/require"`}}, &ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"github.com/stretchr/testify/require"`}},
}, },

View File

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

View File

@@ -3,7 +3,7 @@ package db_test
import ( import (
"fmt" "fmt"
. "git.offline-twitter.com/offline-labs/gas-stack/pkg/flowutils" "git.offline-twitter.com/offline-labs/gas-stack/pkg/must"
. "{{ .ModuleName }}/pkg/db" . "{{ .ModuleName }}/pkg/db"
) )
@@ -14,6 +14,6 @@ func init() {
TestDB = MakeDB("tmp") TestDB = MakeDB("tmp")
} }
func MakeDB(dbName string) *DB { 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 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())
}

61
pkg/schema/column.go Normal file
View File

@@ -0,0 +1,61 @@
package schema
import (
"strings"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/textutils"
)
// Column represents a single column in a table.
type Column struct {
TableName string `db:"table_name"`
Name string `db:"column_name"`
Type string `db:"column_type"`
IsNotNull bool `db:"notnull"`
HasDefaultValue bool `db:"has_default_value"`
DefaultValue string `db:"dflt_value"`
IsPrimaryKey bool `db:"is_primary_key"`
PrimaryKeyRank uint `db:"primary_key_rank"`
IsForeignKey bool `db:"is_foreign_key"`
ForeignKeyTargetTable string `db:"fk_target_table"`
ForeignKeyTargetColumn string `db:"fk_target_column"`
}
// IsNullableForeignKey is a helper function.
func (c Column) IsNullableForeignKey() bool {
return !c.IsNotNull && !c.IsPrimaryKey && c.IsForeignKey
}
func (c Column) IsNonCodeTableForeignKey() bool {
return c.IsForeignKey && strings.HasSuffix(c.Name, "_id")
}
func (c Column) GoFieldName() string {
if c.Name == "rowid" {
return "ID"
}
if c.IsNonCodeTableForeignKey() {
return textutils.SnakeToCamel(strings.TrimSuffix(c.Name, "_id")) + "ID"
}
return textutils.SnakeToCamel(c.Name)
}
// GoVarName returns the name of a local variable for this column, e.g., when used as a function parameter.
func (c Column) GoVarName() string {
if c.Name == "rowid" {
return strings.ToLower(c.TableName)[0:1] + "ID"
// TODO: Or should it just be "id"??
}
// For foreign keys, use first letter of the target type and "ID". "UserID" => "uID"
if c.IsNonCodeTableForeignKey() {
return strings.ToLower(c.ForeignKeyTargetTable)[0:1] + "ID"
}
// Otherwise, just use the whole name
return c.LongGoVarName()
}
// LongGoVarName returns a lowercased version of the field name (Pascal => Camel).
func (c Column) LongGoVarName() string {
return textutils.CamelToPascal(c.GoFieldName())
}

10
pkg/schema/index.go Normal file
View File

@@ -0,0 +1,10 @@
package schema
type Index struct {
Name string `db:"index_name"`
TableName string `db:"table_name"`
Columns []string
IsUnique bool `db:"is_unique"`
// TODO: `where ...` for partial indexes
// TODO: identify columns that are expressions
}

View File

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

View File

@@ -10,7 +10,7 @@ import (
"github.com/jmoiron/sqlx" "github.com/jmoiron/sqlx"
_ "github.com/mattn/go-sqlite3" _ "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" "git.offline-twitter.com/offline-labs/gas-stack/pkg/textutils"
) )
@@ -38,20 +38,20 @@ func SchemaFromDB(db *sqlx.DB) Schema {
db.MustExec(create_views) db.MustExec(create_views)
var tables []Table 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 { for _, tbl := range tables {
tbl.GoTypeName = TypenameFromTablename(tbl.TableName) tbl.GoTypeName = TypenameFromTablename(tbl.TableName)
tbl.TypeIDName = tbl.GoTypeName + "ID" tbl.TypeIDName = tbl.GoTypeName + "ID"
tbl.VarName = strings.ToLower(string(tbl.TableName[0])) 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 ret.Tables[tbl.TableName] = tbl
} }
var indexes []Index 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 { 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 ret.Indexes[idx.Name] = idx
} }
return ret return ret

6
pkg/schema/schema.go Normal file
View File

@@ -0,0 +1,6 @@
package schema
type Schema struct {
Tables map[string]Table
Indexes map[string]Index
}

View File

@@ -2,61 +2,8 @@ package schema
import ( import (
"sort" "sort"
"strings"
"git.offline-twitter.com/offline-labs/gas-stack/pkg/textutils"
) )
// Column represents a single column in a table.
type Column struct {
TableName string `db:"table_name"`
Name string `db:"column_name"`
Type string `db:"column_type"`
IsNotNull bool `db:"notnull"`
HasDefaultValue bool `db:"has_default_value"`
DefaultValue string `db:"dflt_value"`
IsPrimaryKey bool `db:"is_primary_key"`
PrimaryKeyRank uint `db:"primary_key_rank"`
IsForeignKey bool `db:"is_foreign_key"`
ForeignKeyTargetTable string `db:"fk_target_table"`
ForeignKeyTargetColumn string `db:"fk_target_column"`
}
// IsNullableForeignKey is a helper function.
func (c Column) IsNullableForeignKey() bool {
return !c.IsNotNull && !c.IsPrimaryKey && c.IsForeignKey
}
func (c Column) IsNonCodeTableForeignKey() bool {
return c.IsForeignKey && strings.HasSuffix(c.Name, "_id")
}
func (c Column) GoFieldName() string {
if c.Name == "rowid" {
return "ID"
}
if c.IsNonCodeTableForeignKey() {
return textutils.SnakeToCamel(strings.TrimSuffix(c.Name, "_id")) + "ID"
}
return textutils.SnakeToCamel(c.Name)
}
// GoVarName returns the name of a local variable for this column, e.g., when used as a function parameter.
func (c Column) GoVarName() string {
if c.Name == "rowid" {
return strings.ToLower(c.TableName)[0:1] + "ID"
// TODO: Or should it just be "id"??
}
// For foreign keys, use first letter of the target type and "ID". "UserID" => "uID"
if c.IsNonCodeTableForeignKey() {
return strings.ToLower(c.ForeignKeyTargetTable)[0:1] + "ID"
}
// Otherwise, just lowercase the field name
fieldname := c.GoFieldName()
return strings.ToLower(fieldname)[0:1] + fieldname[1:]
}
// Table is a single SQLite table. // Table is a single SQLite table.
type Table struct { type Table struct {
TableName string `db:"name"` TableName string `db:"name"`
@@ -114,17 +61,3 @@ func (t Table) HasAutoTimestamps() (hasCreatedAt bool, hasUpdatedAt bool) {
} }
return return
} }
type Index struct {
Name string `db:"index_name"`
TableName string `db:"table_name"`
Columns []string
IsUnique bool `db:"is_unique"`
// TODO: `where ...` for partial indexes
// TODO: identify columns that are expressions
}
type Schema struct {
Tables map[string]Table
Indexes map[string]Index
}

View File

@@ -9,3 +9,7 @@ func SnakeToCamel(s string) string {
} }
return strings.Join(parts, "") return strings.Join(parts, "")
} }
func CamelToPascal(s string) string {
return strings.ToLower(s)[0:1] + s[1:]
}