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

The test that refuses production calls to (*gorm.DB).Scan, the one GORM path that bypasses the logger's value suppression, overstated what it checks and could pass while skipping a whole package. Its comments now say it matches receiver method names, not types, and name the evasion this leaves; GORM's Rows is dropped from the accepted names. Method values are stated as out of scope with the reason. The file-count floor is replaced by a check that every package the walk parses, static and templates included, was reached. The planted snippets are valid Go and cover each receiver form the guard claims to handle. Test change only.

Model: opus-5-5
This commit was merged in pull request #456.
This commit is contained in:
2026-10-02 19:19:45 +02:00
parent 40f59ec4d2
commit 73353bc8e5
+99 -40
View File
@@ -14,18 +14,16 @@ 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 GORM's Row or database/sql's
// found nothing to look at. The tree held 60 non-test .go files when // QueryRow or QueryRowContext, which return a *sql.Row whose Scan is
// this was written. // database/sql's and not (*gorm.DB).Scan. GORM's Rows is not listed:
const minNonTestFiles = 40 // 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
// isRowProducer reports whether name is a method that returns a // a repo-local method with one of these names that returns *gorm.DB
// database/sql row handle. GORM's Row and Rows return *sql.Row and // gets past it: Scan on that method's result is not reported.
// *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", "QueryRow", "QueryRowContext":
return true return true
default: default:
return false return false
@@ -50,9 +48,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 +114,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 +150,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 +165,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 +197,39 @@ 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)
// 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.GreaterOrEqual(
t, parsed, minNonTestFiles,
"parsed %d non-test .go files, so this check found "+
"nothing to look at", parsed,
)
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 +250,51 @@ type scanGuardCase struct {
want int 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 { 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},
{
"sql QueryRowContext",
`sqlDB.QueryRowContext(ctx, "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 (
"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 // 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 +306,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,