177 lines
4.5 KiB
Go
177 lines
4.5 KiB
Go
package lint_test
|
|
|
|
import (
|
|
"os"
|
|
"path/filepath"
|
|
"testing"
|
|
|
|
"git.warky.dev/wdevs/pgtidy/pkg/diagnostics"
|
|
"git.warky.dev/wdevs/pgtidy/pkg/lint"
|
|
)
|
|
|
|
func fixtureDir() string {
|
|
return filepath.Join("..", "..", "testdata", "lint")
|
|
}
|
|
|
|
func checkFile(t *testing.T, name string) []diagnostics.Diagnostic {
|
|
t.Helper()
|
|
src, err := os.ReadFile(filepath.Join(fixtureDir(), name))
|
|
if err != nil {
|
|
t.Fatalf("read %s: %v", name, err)
|
|
}
|
|
eng := lint.New()
|
|
diags, err := eng.Check(string(src), name)
|
|
if err != nil {
|
|
t.Fatalf("Check(%s): %v", name, err)
|
|
}
|
|
return diags
|
|
}
|
|
|
|
func ruleIDs(diags []diagnostics.Diagnostic) map[string]int {
|
|
m := make(map[string]int)
|
|
for _, d := range diags {
|
|
m[d.RuleID]++
|
|
}
|
|
return m
|
|
}
|
|
|
|
func TestMigrationViolations(t *testing.T) {
|
|
diags := checkFile(t, "migration_violations.sql")
|
|
ids := ruleIDs(diags)
|
|
|
|
cases := []struct {
|
|
id string
|
|
count int
|
|
}{
|
|
{"MIG001", 1}, // one non-concurrent index
|
|
{"MIG002", 1}, // one NOT NULL + no DEFAULT
|
|
{"MIG003", 2}, // one FK + one CHECK without NOT VALID
|
|
}
|
|
for _, c := range cases {
|
|
if got := ids[c.id]; got != c.count {
|
|
t.Errorf("rule %s: want %d findings, got %d", c.id, c.count, got)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestMigrationClean(t *testing.T) {
|
|
diags := checkFile(t, "migration_clean.sql")
|
|
ids := ruleIDs(diags)
|
|
for _, id := range []string{"MIG001", "MIG002", "MIG003"} {
|
|
if n := ids[id]; n != 0 {
|
|
t.Errorf("rule %s: want 0 findings on clean fixture, got %d", id, n)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestCorrectnessViolations(t *testing.T) {
|
|
diags := checkFile(t, "correctness_violations.sql")
|
|
ids := ruleIDs(diags)
|
|
|
|
cases := []struct {
|
|
id string
|
|
count int
|
|
}{
|
|
{"COR001", 1},
|
|
{"COR002", 1},
|
|
{"COR003", 1},
|
|
}
|
|
for _, c := range cases {
|
|
if got := ids[c.id]; got != c.count {
|
|
t.Errorf("rule %s: want %d findings, got %d", c.id, c.count, got)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestNamingViolations(t *testing.T) {
|
|
diags := checkFile(t, "naming_violations.sql")
|
|
ids := ruleIDs(diags)
|
|
|
|
cases := []struct {
|
|
id string
|
|
count int
|
|
}{
|
|
{"NAM001", 1}, // UserAccounts
|
|
{"NAM002", 2}, // userId, emailAddress
|
|
{"NAM003", 1}, // GetUserById
|
|
}
|
|
for _, c := range cases {
|
|
if got := ids[c.id]; got != c.count {
|
|
t.Errorf("rule %s: want %d findings, got %d", c.id, c.count, got)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestParseError(t *testing.T) {
|
|
eng := lint.New()
|
|
diags, err := eng.Check("SELECT FROM WHERE", "test.sql")
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
if len(diags) == 0 || diags[0].RuleID != "PARSE" {
|
|
t.Errorf("expected a PARSE diagnostic, got %v", diags)
|
|
}
|
|
}
|
|
|
|
func TestCustomEngine(t *testing.T) {
|
|
eng := lint.Custom()
|
|
eng.Register(lint.New().Rules()[0]) // register first rule only
|
|
// Just confirm it doesn't panic and returns results.
|
|
diags, err := eng.Check("SELECT 1", "stdin")
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
_ = diags
|
|
}
|
|
|
|
func TestPlpgsqlBodyLint(t *testing.T) {
|
|
src := "create function f() returns void language plpgsql as $$\nbegin\n perform 1;\n select * from t;\nend;\n$$;\n" +
|
|
"create function g() returns void language plpgsql as $$\nbegin\n selec 1;\nend;\n$$;\n" +
|
|
"create function p() returns void language plpython3u as $$\nselect * from nothing\n$$;\n" +
|
|
"do $$ begin select * from t; end $$;\n"
|
|
diags, err := lint.New().Check(src, "x.sql")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var cor001, pl []diagnostics.Diagnostic
|
|
for _, d := range diags {
|
|
switch d.RuleID {
|
|
case "COR001":
|
|
cor001 = append(cor001, d)
|
|
case "PLPGSQL":
|
|
pl = append(pl, d)
|
|
}
|
|
}
|
|
if len(cor001) != 2 || cor001[0].Line != 4 || cor001[0].Col != 10 || cor001[1].Line != 15 {
|
|
t.Errorf("COR001 in bodies: %+v", cor001)
|
|
}
|
|
for _, d := range cor001 {
|
|
if d.Fix != nil {
|
|
t.Errorf("embedded diagnostic must not carry fixes: %+v", d)
|
|
}
|
|
}
|
|
if len(pl) != 1 || pl[0].Line != 9 || pl[0].Col != 3 {
|
|
t.Errorf("PLPGSQL syntax error: %+v", pl)
|
|
}
|
|
}
|
|
|
|
func TestEmbeddedDollarSQLLint(t *testing.T) {
|
|
src := "create function f() returns void language plpgsql as $$\nbegin\n execute $q$\n select * from t\n $q$;\nend;\n$$;\n" +
|
|
"create function s() returns int language sql as $$ select * from u $$;\n" +
|
|
"create function p() returns void language plpython3u as $$\nselect * from nothing\n$$;\n" +
|
|
"select format($f$select %I from t$f$, 'a');\n"
|
|
diags, err := lint.New().Check(src, "x.sql")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
var got [][2]int
|
|
for _, d := range diags {
|
|
if d.RuleID == "COR001" {
|
|
got = append(got, [2]int{d.Line, d.Col})
|
|
}
|
|
}
|
|
if len(got) != 2 || got[0] != [2]int{4, 12} || got[1] != [2]int{8, 59} {
|
|
t.Errorf("COR001 positions: %v", got)
|
|
}
|
|
}
|