Files
webhooker/internal/gormlog/scan_guard_test.go
T
sneak 9d963cac83
check / check (push) Successful in 3m23s
Harden the (*gorm.DB).Scan guard test (closes #232)
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
2026-10-02 16:16:45 +00:00

307 lines
7.3 KiB
Go

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 one of the GORM and
// database/sql methods that return a row handle, whose Scan is
// database/sql's and not (*gorm.DB).Scan. 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", "Rows", "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)
// Every directory directly under cmd and internal is a package with
// non-test code, so one the walk parsed nothing in was skipped, and
// a Scan there would pass unseen.
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(
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 <producer>().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
// the row producers it 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},
{"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
// 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)
})
}
}