package main import ( "errors" "fmt" "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/must" "git.offline-twitter.com/offline-labs/gas-stack/pkg/schema" ) var ErrNoSuchTable = errors.New("no such table") var generate_model = &cobra.Command{ Use: "generate ", Short: "Generate a model type", Args: cobra.ExactArgs(1), RunE: func(cmd *cobra.Command, args []string) error { 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)) 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, sch, modname) must.Do(modelgenerate.FprintWithComments(os.Stdout, file2)) } else { decls := []ast.Decl{ &ast.GenDecl{ Tok: token.IMPORT, Specs: []ast.Spec{ &ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"database/sql"`}}, &ast.ImportSpec{Path: &ast.BasicLit{Kind: token.STRING, Value: `"errors"`}}, &ast.ImportSpec{ Name: ast.NewIdent("."), Path: &ast.BasicLit{Kind: token.STRING, Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/db"`}, }, &ast.ImportSpec{ Path: &ast.BasicLit{ Kind: token.STRING, Value: `"git.offline-twitter.com/offline-labs/gas-stack/pkg/must"`, }, }, }, }, } if !table.IsWithoutRowid { decls = append(decls, modelgenerate.GenerateIDType(table)) } decls = append(decls, modelgenerate.GenerateModelAST(table), modelgenerate.GenerateSQLFieldsConst(table), modelgenerate.GenerateSaveItemFunc(table), modelgenerate.GenerateDeleteItemFunc(table), ) if table.IsWithoutRowid { decls = append(decls, modelgenerate.GenerateGetItemBy(table, table.PrimaryKeyColumns()), ) } else { decls = append(decls, modelgenerate.GenerateGetItemByIDFunc(table), ) } for _, index := range sch.Indexes { if index.TableName != table.TableName { // Skip indexes on other tables continue } 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)) file := &ast.File{ Name: ast.NewIdent("db"), // TODO: parameterize Decls: decls, } must.Do(modelgenerate.FprintWithComments(os.Stdout, file)) } return nil }, } func init() { generate_model.Flags().String("schema", "pkg/db/schema.sql", "Path to SQL schema file") generate_model.Flags().String("modname", "mymodule", "Name of project's Go module (TODO: detect automatically)") generate_model.Flags().Bool("test", false, "Generate test file instead of regular file") }