package gormlog_test import ( "fmt" "go/ast" "go/parser" "go/token" "io/fs" "os" "path/filepath" "strings" "testing" "github.com/stretchr/testify/require" ) // isRowProducer reports whether name is GORM's Row or database/sql's // QueryRow or QueryRowContext, which return a *sql.Row whose Scan is // database/sql's and not (*gorm.DB).Scan. GORM's Rows is not listed: // it also returns an error, so Scan is never called on its result // directly. It matches the method name only and resolves no types, so // a repo-local method with one of these names that returns *gorm.DB // gets past it: Scan on that method's result is not reported. func isRowProducer(name string) bool { switch name { case "Row", "QueryRow", "QueryRowContext": return true default: return false } } // receiverIsRowHandle reports whether x is syntactically a call to a // row producer, which is the only receiver form this check accepts for // a Scan. func receiverIsRowHandle(x ast.Expr) bool { call, ok := x.(*ast.CallExpr) if !ok { return false } sel, ok := call.Fun.(*ast.SelectorExpr) if !ok { return false } return isRowProducer(sel.Sel.Name) } // unguardedScans returns the position of every Scan call in file whose // receiver is not a call to a row producer. It fails closed: any other // receiver — a local variable, a struct field, a call to any other // method — is reported rather than assumed safe. // // It sees only calls written x.Scan(...). A method value, f := db.Scan // followed by f(&v), is out of scope: Scan is never the called // expression there, and nobody writes a query that way by accident, // which is the mistake this check exists to catch. func unguardedScans( fset *token.FileSet, file *ast.File, ) []token.Position { var found []token.Position ast.Inspect(file, func(n ast.Node) bool { call, ok := n.(*ast.CallExpr) if !ok { return true } sel, ok := call.Fun.(*ast.SelectorExpr) if !ok || sel.Sel.Name != "Scan" { return true } if !receiverIsRowHandle(sel.X) { found = append(found, fset.Position(sel.Sel.Pos())) } return true }) return found } // moduleRoot walks up from the working directory to the directory // holding go.mod. func moduleRoot(t *testing.T) string { t.Helper() dir, err := os.Getwd() require.NoError(t, err) for { _, statErr := os.Stat(filepath.Join(dir, "go.mod")) if statErr == nil { return dir } parent := filepath.Dir(dir) require.NotEqual(t, parent, dir, "no go.mod above %s", dir) dir = parent } } // skipDir reports whether a directory holds no source this check // governs. func skipDir(name string) bool { switch name { case ".git", "bin", "node_modules", "testdata": return true default: return false } } // walkNonTestGo parses every non-test .go file under root. It returns // the directories, relative to root, it parsed a file in, along with // every unguarded Scan it found. func walkNonTestGo(t *testing.T, root string) (map[string]bool, []string) { t.Helper() walked := map[string]bool{} var hits []string fset := token.NewFileSet() require.NoError(t, filepath.WalkDir( root, func(path string, d fs.DirEntry, err error) error { if err != nil { return err } if d.IsDir() { if skipDir(d.Name()) { return fs.SkipDir } return nil } if !isNonTestGo(d.Name()) { return nil } file, err := parser.ParseFile(fset, path, nil, 0) if err != nil { return err } dir, err := filepath.Rel(root, filepath.Dir(path)) if err != nil { return err } walked[dir] = true for _, pos := range unguardedScans(fset, file) { hits = append(hits, relPosition(root, pos)) } return nil }, )) return walked, hits } // isNonTestGo reports whether a file name is Go source this check // governs. func isNonTestGo(name string) bool { return strings.HasSuffix(name, ".go") && !strings.HasSuffix(name, "_test.go") } // relPosition renders pos with its path relative to root, so a failure // names the file the way the repository does. func relPosition(root string, pos token.Position) string { name := pos.Filename rel, err := filepath.Rel(root, name) if err == nil { name = rel } return fmt.Sprintf("%s:%d:%d", name, pos.Line, pos.Column) } // TestGormScanIsNeverCalledOutsideTests keeps (*gorm.DB).Scan out of // non-test code. // // It is the one statement path (*Logger).ParamsFilter does not reach: // Scan swaps GORM's own trace recorder in for the adapter, and that // recorder does not implement gorm.ParamsFilter, so the statement is // logged with its values interpolated. The package comment states the // limit; this fails when someone adds a call site anyway. // // Test files are not governed: what a test binds is fixture data. func TestGormScanIsNeverCalledOutsideTests(t *testing.T) { t.Parallel() root := moduleRoot(t) walked, offenders := walkNonTestGo(t, root) // The module's packages are static, templates, and every directory // directly under cmd and internal. Each holds non-test code, so one // the walk parsed nothing in was skipped, and a Scan there would // pass unseen. packages := []string{"static", "templates"} for _, parent := range []string{"cmd", "internal"} { entries, err := os.ReadDir(filepath.Join(root, parent)) require.NoError(t, err) for _, entry := range entries { if !entry.IsDir() { continue } packages = append(packages, filepath.Join(parent, entry.Name())) } } for _, dir := range packages { require.True( t, walked[dir], "the walk parsed no non-test .go file in %s", dir, ) } require.Empty( t, offenders, "Scan called on a receiver this check cannot show is a "+ "database/sql row handle. (*gorm.DB).Scan logs the "+ "statement with its bound values interpolated — use "+ "Find, Pluck, or Raw(...).Row().Scan instead. A "+ "database/sql Scan reached through a variable is "+ "reported too; write it as ().Scan rather "+ "than widening this check.", ) } // scanGuardCase is one planted snippet and whether the check above // should report it. type scanGuardCase struct { name string body string want int } // scanGuardCases covers each receiver form unguardedScans names, plus // each row producer isRowProducer lets through. Each body is valid Go // inside plantedFile. func scanGuardCases() []scanGuardCase { return []scanGuardCase{ {"local variable", "q := gdb.Raw(\"SELECT 1\")\n\tq.Scan(&v)", 1}, {"struct field", `s.db.Scan(&v)`, 1}, {"gorm chain", `gdb.Raw("SELECT 1").Scan(&v)`, 1}, { "sql rows in a variable", "rows, _ := gdb.Raw(\"SELECT 1\").Rows()\n\trows.Scan(&v)", 1, }, {"gorm Row", `gdb.Raw("SELECT 1").Row().Scan(&v)`, 0}, {"sql QueryRow", `sqlDB.QueryRow("SELECT 1").Scan(&v)`, 0}, { "sql QueryRowContext", `sqlDB.QueryRowContext(ctx, "SELECT 1").Scan(&v)`, 0, }, {"unrelated call", `gdb.Find(&v)`, 0}, } } // plantedFile wraps one case body in a function that declares every // name the bodies use, so each body is the Go it stands for. The result // is parsed, never compiled. const plantedFile = `package p import ( "context" "database/sql" "gorm.io/gorm" ) type store struct{ db *gorm.DB } func f(ctx context.Context, gdb *gorm.DB, sqlDB *sql.DB, s store) { var v int %s } ` // TestScanGuard_ReportsPlantedCalls proves the check fires. Without it // a detector that matched nothing would satisfy the walk above no // matter what the tree contained. func TestScanGuard_ReportsPlantedCalls(t *testing.T) { t.Parallel() for _, tc := range scanGuardCases() { t.Run(tc.name, func(t *testing.T) { t.Parallel() fset := token.NewFileSet() src := fmt.Sprintf(plantedFile, tc.body) file, err := parser.ParseFile( fset, tc.name+".go", src, 0, ) require.NoError(t, err) require.Len(t, unguardedScans(fset, file), tc.want) }) } }