Harden the (*gorm.DB).Scan guard test (closes #232)
check / check (push) Successful in 3m23s

The isRowProducer comment now says it matches method names only and
names the evasion that follows: a repo-local Row or QueryRow helper
returning *gorm.DB gets past it. unguardedScans states method values
(f := db.Scan) as out of scope, with the reason. The 40-file floor is
replaced by a check that every directory under cmd and internal was
walked, so skipping a whole package fails the test. The planted
snippets now cover the struct-field receiver and are all valid Go
inside a wrapper that declares the names they use; the Rows snippet
that could not compile is replaced by the variable form the guard
reports. The stale "one caller" sentence is dropped.

Model: opus-5-5
This commit is contained in:
2026-10-02 16:16:45 +00:00
parent 45bd7e9b94
commit 9d963cac83
+84 -38
View File
@@ -14,15 +14,12 @@ import (
"github.com/stretchr/testify/require" "github.com/stretchr/testify/require"
) )
// minNonTestFiles guards the walk below against passing because it // isRowProducer reports whether name is one of the GORM and
// found nothing to look at. The tree held 60 non-test .go files when // database/sql methods that return a row handle, whose Scan is
// this was written. // database/sql's and not (*gorm.DB).Scan. It matches the method name
const minNonTestFiles = 40 // 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
// isRowProducer reports whether name is a method that returns a // result is not reported.
// database/sql row handle. GORM's Row and Rows return *sql.Row and
// *sql.Rows, so Scan on the result of one of them is database/sql's
// Scan and never (*gorm.DB).Scan.
func isRowProducer(name string) bool { func isRowProducer(name string) bool {
switch name { switch name {
case "Row", "Rows", "QueryRow", "QueryRowContext": case "Row", "Rows", "QueryRow", "QueryRowContext":
@@ -50,9 +47,14 @@ func receiverIsRowHandle(x ast.Expr) bool {
} }
// unguardedScans returns the position of every Scan call in file whose // unguardedScans returns the position of every Scan call in file whose
// receiver is not a row handle. It fails closed: a receiver it cannot // receiver is not a call to a row producer. It fails closed: any other
// resolve syntactically — a local variable, a struct field — is // receiver — a local variable, a struct field, a call to any other
// reported rather than assumed safe. // 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( func unguardedScans(
fset *token.FileSet, file *ast.File, fset *token.FileSet, file *ast.File,
) []token.Position { ) []token.Position {
@@ -111,15 +113,15 @@ func skipDir(name string) bool {
} }
} }
// walkNonTestGo parses every non-test .go file under root and returns // walkNonTestGo parses every non-test .go file under root. It returns
// how many it parsed along with every unguarded Scan it found. // the directories, relative to root, it parsed a file in, along with
func walkNonTestGo(t *testing.T, root string) (int, []string) { // every unguarded Scan it found.
func walkNonTestGo(t *testing.T, root string) (map[string]bool, []string) {
t.Helper() t.Helper()
var ( walked := map[string]bool{}
parsed int
hits []string var hits []string
)
fset := token.NewFileSet() fset := token.NewFileSet()
@@ -147,7 +149,12 @@ func walkNonTestGo(t *testing.T, root string) (int, []string) {
return err return err
} }
parsed++ dir, err := filepath.Rel(root, filepath.Dir(path))
if err != nil {
return err
}
walked[dir] = true
for _, pos := range unguardedScans(fset, file) { for _, pos := range unguardedScans(fset, file) {
hits = append(hits, relPosition(root, pos)) hits = append(hits, relPosition(root, pos))
@@ -157,7 +164,7 @@ func walkNonTestGo(t *testing.T, root string) (int, []string) {
}, },
)) ))
return parsed, hits return walked, hits
} }
// isNonTestGo reports whether a file name is Go source this check // isNonTestGo reports whether a file name is Go source this check
@@ -189,19 +196,33 @@ func relPosition(root string, pos token.Position) string {
// logged with its values interpolated. The package comment states the // logged with its values interpolated. The package comment states the
// limit; this fails when someone adds a call site anyway. // limit; this fails when someone adds a call site anyway.
// //
// The current tree has one caller, internal/database/database_test.go, // Test files are not governed: what a test binds is fixture data.
// which this check does not govern: it is test-only and its SELECT 1
// binds nothing.
func TestGormScanIsNeverCalledOutsideTests(t *testing.T) { func TestGormScanIsNeverCalledOutsideTests(t *testing.T) {
t.Parallel() t.Parallel()
parsed, offenders := walkNonTestGo(t, moduleRoot(t)) root := moduleRoot(t)
walked, offenders := walkNonTestGo(t, root)
require.GreaterOrEqual( // Every directory directly under cmd and internal is a package with
t, parsed, minNonTestFiles, // non-test code, so one the walk parsed nothing in was skipped, and
"parsed %d non-test .go files, so this check found "+ // a Scan there would pass unseen.
"nothing to look at", parsed, 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
}
dir := filepath.Join(parent, entry.Name())
require.True(
t, walked[dir],
"the walk parsed no non-test .go file in %s", dir,
) )
}
}
require.Empty( require.Empty(
t, offenders, t, offenders,
"Scan called on a receiver this check cannot show is a "+ "Scan called on a receiver this check cannot show is a "+
@@ -222,18 +243,45 @@ type scanGuardCase struct {
want int want int
} }
// scanGuardCases covers each receiver form unguardedScans names, plus
// the row producers it lets through. Each body is valid Go inside
// plantedFile.
func scanGuardCases() []scanGuardCase { func scanGuardCases() []scanGuardCase {
return []scanGuardCase{ return []scanGuardCase{
{"gorm chain", `db.DB().Raw("SELECT 1").Scan(&v)`, 1}, {"local variable", "q := gdb.Raw(\"SELECT 1\")\n\tq.Scan(&v)", 1},
{"gorm receiver", `gdb.Scan(&v)`, 1}, {"struct field", `s.db.Scan(&v)`, 1},
{"gorm via variable", "q := gdb.Raw(\"x\")\nq.Scan(&v)", 1}, {"gorm chain", `gdb.Raw("SELECT 1").Scan(&v)`, 1},
{"gorm model chain", `gdb.Model(&x).Scan(&v)`, 1}, {
{"sql row", `gdb.Raw("SELECT 1").Row().Scan(&v)`, 0}, "sql rows in a variable",
{"sql rows", `gdb.Raw("SELECT 1").Rows().Scan(&v)`, 0}, "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},
{"unrelated call", `gdb.Find(&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 (
"database/sql"
"gorm.io/gorm"
)
type store struct{ db *gorm.DB }
func f(gdb *gorm.DB, sqlDB *sql.DB, s store) {
var v int
%s
}
`
// TestScanGuard_ReportsPlantedCalls proves the check fires. Without it // TestScanGuard_ReportsPlantedCalls proves the check fires. Without it
// a detector that matched nothing would satisfy the walk above no // a detector that matched nothing would satisfy the walk above no
// matter what the tree contained. // matter what the tree contained.
@@ -245,9 +293,7 @@ func TestScanGuard_ReportsPlantedCalls(t *testing.T) {
t.Parallel() t.Parallel()
fset := token.NewFileSet() fset := token.NewFileSet()
src := fmt.Sprintf( src := fmt.Sprintf(plantedFile, tc.body)
"package p\n\nfunc f() {\n\t%s\n}\n", tc.body,
)
file, err := parser.ParseFile( file, err := parser.ParseFile(
fset, tc.name+".go", src, 0, fset, tc.name+".go", src, 0,