diff --git a/cmd/subcmd_generate_models.go b/cmd/subcmd_generate_models.go index 2b9f10d..3fd1d95 100644 --- a/cmd/subcmd_generate_models.go +++ b/cmd/subcmd_generate_models.go @@ -6,6 +6,7 @@ import ( "go/ast" "go/token" "os" + "slices" "github.com/spf13/cobra" @@ -30,14 +31,14 @@ var generate_model = &cobra.Command{ 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.Get(cmd.Flags().GetBool("test")) { - file2 := modelgenerate.GenerateModelTestAST(table, schema, modname) + file2 := modelgenerate.GenerateModelTestAST(table, sch, modname) must.Do(modelgenerate.FprintWithComments(os.Stdout, file2)) } else { decls := []ast.Decl{ @@ -77,19 +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 len(index.Columns) != 1 || index.Columns[0] == "" { - // Skip multi-column and expression indexes + 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.GenerateGetItemByUniqColFunc(table, table.GetColumnByName(index.Columns[0]))) + decls = append(decls, modelgenerate.GenerateGetItemBy(table, cols)) } else { - decls = append(decls, modelgenerate.GenerateGetItemsByColFunc(table, table.GetColumnByName(index.Columns[0]))) + decls = append(decls, modelgenerate.GenerateGetItemsBy(table, cols)) } } decls = append(decls, modelgenerate.GenerateGetAllItemsFunc(table)) diff --git a/pkg/codegen/modelgenerate/generate_testfile.go b/pkg/codegen/modelgenerate/generate_testfile.go index 314411e..8be63e4 100644 --- a/pkg/codegen/modelgenerate/generate_testfile.go +++ b/pkg/codegen/modelgenerate/generate_testfile.go @@ -4,6 +4,7 @@ import ( "fmt" "go/ast" "go/token" + "slices" "strings" "github.com/jinzhu/inflection" @@ -432,11 +433,20 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam // Skip indexes on other tables continue } - if len(index.Columns) != 1 || index.Columns[0] == "" { - // Skip multi-column and expression indexes + if slices.Contains(index.Columns, "") { + // Skip expression indexes; there's no way to resolve an expression to a real column continue } - col := tbl.GetColumnByName(index.Columns[0]) + 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(...)) @@ -447,9 +457,9 @@ 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, }), }, }}, @@ -463,9 +473,9 @@ func GenerateModelTestAST(tbl pkgschema.Table, schema pkgschema.Schema, gomodNam ast.NewIdent("t"), &ast.CallExpr{ Fun: &ast.SelectorExpr{X: testDB, Sel: ast.NewIdent( - "Get" + inflection.Plural(pkgschema.TypenameFromTablename(tbl.TableName)) + "By" + col.GoFieldName(), + "Get" + inflection.Plural(pkgschema.TypenameFromTablename(tbl.TableName)) + "By" + funcNameSuffix, )}, - Args: []ast.Expr{&ast.SelectorExpr{X: testObj2, Sel: ast.NewIdent(col.GoFieldName())}}, + Args: callArgs, }, testObj2, },