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) } }