Harden the (*gorm.DB).Scan guard test (closes #232)
check / check (push) Successful in 3m25s
check / check (push) Successful in 3m25s
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. GORM's Rows is dropped from those names: it also returns an error, so Scan is never called on its result directly. unguardedScans states method values (f := db.Scan) as out of scope, with the reason. The 40-file floor is replaced by a check that static, templates and every directory under cmd and internal was walked, so skipping a whole package fails the test. The planted snippets cover the struct-field receiver and each accepted row producer, and are all valid Go inside a wrapper that declares the names they use. The stale "one caller" sentence is dropped. Model: opus-5-5
This commit is contained in:
@@ -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)
|
||||||
|
|
||||||
require.GreaterOrEqual(
|
// The module's packages are static, templates, and every directory
|
||||||
t, parsed, minNonTestFiles,
|
// directly under cmd and internal. Each holds non-test code, so one
|
||||||
"parsed %d non-test .go files, so this check found "+
|
// the walk parsed nothing in was skipped, and a Scan there would
|
||||||
"nothing to look at", parsed,
|
// 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(
|
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,
|
||||||
|
|||||||
Reference in New Issue
Block a user