package build import ( "bytes" "fmt" "go/ast" "go/parser" "go/token" "os" "path/filepath" "sort" "strings" ) const ( registryFileName = "registry.gen.go" generatedMark = "Code generated by summer make" ) type identScanKind int const ( scanStructs identScanKind = iota scanMigrationFuncs scanCommandFuncs scanJobFuncs scanAdminFuncs ) type registryRef struct { Package string Ident string SortKey string } type registryData struct { Package string Module string Imports []string Models []registryRef Migrations []registryRef Commands []registryRef Jobs []registryRef AdminControllers []registryRef } func emptyRegistry(pkg, module string) registryData { data := registryData{Package: pkg, Module: module} data.prepare() return data } func (d *registryData) prepare() { d.Imports = uniqueLeafImports(d.Module, d.Models, d.Migrations, d.Commands, d.Jobs, d.AdminControllers) } func uniqueLeafImports(module string, groups ...[]registryRef) []string { if module == "" { return nil } seen := make(map[string]struct{}) var imports []string for _, group := range groups { for _, ref := range group { if ref.Package == "" { continue } path := module + "/" + ref.Package if _, ok := seen[path]; ok { continue } seen[path] = struct{}{} imports = append(imports, path) } } sort.Strings(imports) return imports } func refreshRegistry(pluginDir string) error { pkg, err := packageNameOf(filepath.Join(pluginDir, "plugin.go")) if err != nil { return err } module, err := readModulePath(filepath.Join(pluginDir, "go.mod")) if err != nil { return err } data := registryData{Package: pkg, Module: module} if data.Models, err = scanGeneratedIdents(filepath.Join(pluginDir, "models"), "models", scanStructs); err != nil { return err } if data.Migrations, err = scanGeneratedIdents(filepath.Join(pluginDir, "updates"), "updates", scanMigrationFuncs); err != nil { return err } if data.Commands, err = scanGeneratedIdents(filepath.Join(pluginDir, "console"), "console", scanCommandFuncs); err != nil { return err } if data.Jobs, err = scanGeneratedIdents(filepath.Join(pluginDir, "jobs"), "jobs", scanJobFuncs); err != nil { return err } if data.AdminControllers, err = scanGeneratedIdents(filepath.Join(pluginDir, "controllers"), "controllers", scanAdminFuncs); err != nil { return err } sortRefs(data.Models) sortRefs(data.Migrations) sortRefs(data.Commands) sortRefs(data.Jobs) sortRefs(data.AdminControllers) return writeRegistry(pluginDir, data) } func packageNameOf(path string) (string, error) { fset := token.NewFileSet() file, err := parser.ParseFile(fset, path, nil, parser.PackageClauseOnly) if err != nil { return "", fmt.Errorf("build: parse %s: %w", path, err) } if file.Name == nil || file.Name.Name == "" { return "", fmt.Errorf("build: %s has no package name", path) } return file.Name.Name, nil } func scanGeneratedIdents(dir, pkg string, kind identScanKind) ([]registryRef, error) { entries, err := os.ReadDir(dir) if err != nil { if os.IsNotExist(err) { return nil, nil } return nil, fmt.Errorf("build: read %s: %w", dir, err) } fset := token.NewFileSet() var refs []registryRef for _, entry := range entries { name := entry.Name() if entry.IsDir() || !strings.HasSuffix(name, ".go") || strings.HasSuffix(name, "_test.go") || name == "doc.go" { continue } path := filepath.Join(dir, name) file, err := parser.ParseFile(fset, path, nil, parser.ParseComments) if err != nil { return nil, fmt.Errorf("build: parse %s: %w", path, err) } if !isGeneratedFile(file) { continue } switch kind { case scanStructs: refs = append(refs, structRefs(pkg, name, file)...) default: want := resultTypeFor(kind) refs = append(refs, funcRefs(pkg, name, file, want)...) } } return refs, nil } func resultTypeFor(kind identScanKind) string { switch kind { case scanMigrationFuncs: return "gormigrate.Migration" case scanCommandFuncs: return "bonfire.Command" case scanJobFuncs: return "pact.Job" case scanAdminFuncs: return "pact.AdminController" default: return "" } } func isGeneratedFile(file *ast.File) bool { for _, cg := range file.Comments { for _, c := range cg.List { if strings.Contains(c.Text, generatedMark) { return true } } } return false } func structRefs(pkg, filename string, file *ast.File) []registryRef { var refs []registryRef for _, decl := range file.Decls { gd, ok := decl.(*ast.GenDecl) if !ok { continue } for _, spec := range gd.Specs { ts, ok := spec.(*ast.TypeSpec) if !ok || ts.Name == nil || !ts.Name.IsExported() { continue } if _, ok := ts.Type.(*ast.StructType); !ok { continue } refs = append(refs, registryRef{Package: pkg, Ident: ts.Name.Name, SortKey: ts.Name.Name}) } } _ = filename return refs } func funcRefs(pkg, filename string, file *ast.File, want string) []registryRef { var refs []registryRef for _, decl := range file.Decls { fn, ok := decl.(*ast.FuncDecl) if !ok || fn.Recv != nil || fn.Name == nil || !fn.Name.IsExported() || fn.Type == nil { continue } if fn.Type.Params != nil && len(fn.Type.Params.List) > 0 { continue } if fn.Type.Results == nil || len(fn.Type.Results.List) != 1 { continue } if astTypeName(fn.Type.Results.List[0].Type) != want { continue } refs = append(refs, registryRef{ Package: pkg, Ident: fn.Name.Name, SortKey: strings.TrimSuffix(filename, ".go"), }) } return refs } func astTypeName(expr ast.Expr) string { switch t := expr.(type) { case *ast.StarExpr: return astTypeName(t.X) case *ast.SelectorExpr: if id, ok := t.X.(*ast.Ident); ok && t.Sel != nil { return id.Name + "." + t.Sel.Name } case *ast.Ident: return t.Name } return "" } func sortRefs(refs []registryRef) { sort.Slice(refs, func(i, j int) bool { a, b := refs[i].SortKey, refs[j].SortKey if a == "" { a = refs[i].Ident } if b == "" { b = refs[j].Ident } if a == b { return refs[i].Ident < refs[j].Ident } return a < b }) } func writeRegistry(pluginDir string, data registryData) error { data.prepare() src, err := renderStub("registry.go", data) if err != nil { return err } path := filepath.Join(pluginDir, registryFileName) existing, err := os.ReadFile(path) if err == nil && bytes.Equal(existing, src) { return nil } if err != nil && !os.IsNotExist(err) { return fmt.Errorf("build: read %s: %w", path, err) } tmp, err := os.CreateTemp(pluginDir, registryFileName+".*") if err != nil { return fmt.Errorf("build: create %s temp: %w", registryFileName, err) } tmpName := tmp.Name() if _, err := tmp.Write(src); err != nil { tmp.Close() os.Remove(tmpName) return fmt.Errorf("build: write %s temp: %w", registryFileName, err) } if err := tmp.Close(); err != nil { os.Remove(tmpName) return fmt.Errorf("build: close %s temp: %w", registryFileName, err) } if err := os.Rename(tmpName, path); err != nil { os.Remove(tmpName) return fmt.Errorf("build: rename %s: %w", registryFileName, err) } return nil }