test: expand coverage across readers, writers, cmd, ui, diff and merge
Implements tests/_plans and previously deferred packages; updates plan README with new coverage numbers.
This commit is contained in:
@@ -0,0 +1,110 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func TestCurrentColumnHasDescription(t *testing.T) {
|
||||
table := models.InitTable("users", "public")
|
||||
c := models.InitColumn("Email", "users", "public")
|
||||
c.Description = " the email "
|
||||
table.Columns["Email"] = c
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
table *models.Table
|
||||
col *models.Column
|
||||
want bool
|
||||
}{
|
||||
{"nil table", nil, &models.Column{Name: "email", Description: "x"}, false},
|
||||
{"match ignoring case and whitespace", table, &models.Column{Name: "email", Description: "the email"}, true},
|
||||
{"different description", table, &models.Column{Name: "email", Description: "other"}, false},
|
||||
{"column missing", table, &models.Column{Name: "age", Description: "x"}, false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := currentColumnHasDescription(tt.table, tt.col); got != tt.want {
|
||||
t.Errorf("got %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecuteCommentColumn(t *testing.T) {
|
||||
te, err := NewTemplateExecutor(false)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got, err := te.ExecuteCommentColumn(CommentColumnData{
|
||||
SchemaName: "public", TableName: "users", ColumnName: "email", Comment: "it''s",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.Contains(got, "COMMENT ON COLUMN") || !strings.Contains(got, "public.users") ||
|
||||
!strings.Contains(got, "email") || !strings.Contains(got, "IS 'it''s';") {
|
||||
t.Errorf("unexpected output: %s", got)
|
||||
}
|
||||
}
|
||||
|
||||
func migrationWithColumnDescription(t *testing.T, currentDesc string, withCurrentCol bool) string {
|
||||
t.Helper()
|
||||
newDB := func(desc string, include bool) *models.Database {
|
||||
db := models.InitDatabase("testdb")
|
||||
s := models.InitSchema("public")
|
||||
tbl := models.InitTable("users", "public")
|
||||
id := models.InitColumn("id", "users", "public")
|
||||
id.Type = "integer"
|
||||
tbl.Columns["id"] = id
|
||||
if include {
|
||||
col := models.InitColumn("email", "users", "public")
|
||||
col.Type = "text"
|
||||
col.Description = desc
|
||||
tbl.Columns["email"] = col
|
||||
}
|
||||
s.Tables = append(s.Tables, tbl)
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
return db
|
||||
}
|
||||
model := newDB("it's the email", true)
|
||||
current := newDB(currentDesc, withCurrentCol)
|
||||
|
||||
var buf bytes.Buffer
|
||||
w, err := NewMigrationWriter(&writers.WriterOptions{})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
w.writer = &buf
|
||||
if err := w.WriteMigration(model, current); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return buf.String()
|
||||
}
|
||||
|
||||
func TestWriteMigration_ColumnComments(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
currentDesc string
|
||||
withCol bool
|
||||
wantComment bool
|
||||
}{
|
||||
{"added", "", true, true},
|
||||
{"changed", "old text", true, true},
|
||||
{"unchanged", "it's the email", true, false},
|
||||
{"new column", "", false, true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
out := migrationWithColumnDescription(t, tt.currentDesc, tt.withCol)
|
||||
has := strings.Contains(out, "COMMENT ON COLUMN") && strings.Contains(out, "it''s the email")
|
||||
if has != tt.wantComment {
|
||||
t.Errorf("comment emitted = %v, want %v\n%s", has, tt.wantComment, out)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,243 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/jackc/pgx/v5"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func liveWriterConn(t *testing.T) string {
|
||||
t.Helper()
|
||||
conn := os.Getenv("RELSPEC_TEST_PG_CONN")
|
||||
if conn == "" {
|
||||
t.Skip("RELSPEC_TEST_PG_CONN not set")
|
||||
}
|
||||
return conn
|
||||
}
|
||||
|
||||
// liveWriterSchema returns a unique schema name and drops it on cleanup.
|
||||
func liveWriterSchema(t *testing.T, connString string) (string, *pgx.Conn) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
conn, err := pgx.Connect(ctx, connString)
|
||||
if err != nil {
|
||||
t.Fatalf("connect: %v", err)
|
||||
}
|
||||
name := fmt.Sprintf("pgw_test_%d", time.Now().UnixNano())
|
||||
t.Cleanup(func() {
|
||||
_, _ = conn.Exec(ctx, "DROP SCHEMA IF EXISTS "+name+" CASCADE")
|
||||
_ = conn.Close(ctx)
|
||||
})
|
||||
return name, conn
|
||||
}
|
||||
|
||||
func liveModel(schemaName string, columns map[string]string) *models.Database {
|
||||
db := models.InitDatabase("live")
|
||||
s := models.InitSchema(schemaName)
|
||||
tbl := models.InitTable("accounts", schemaName)
|
||||
id := models.InitColumn("id", "accounts", schemaName)
|
||||
id.Type = "integer"
|
||||
id.NotNull = true
|
||||
id.IsPrimaryKey = true
|
||||
tbl.Columns["id"] = id
|
||||
for name, typ := range columns {
|
||||
c := models.InitColumn(name, "accounts", schemaName)
|
||||
c.Type = typ
|
||||
tbl.Columns[name] = c
|
||||
}
|
||||
s.Tables = append(s.Tables, tbl)
|
||||
db.Schemas = append(db.Schemas, s)
|
||||
return db
|
||||
}
|
||||
|
||||
func runLiveWrite(t *testing.T, connString string, db *models.Database, meta map[string]interface{}) (*ExecutionReport, error) {
|
||||
t.Helper()
|
||||
m := map[string]interface{}{"connection_string": connString}
|
||||
for k, v := range meta {
|
||||
m[k] = v
|
||||
}
|
||||
w := NewWriter(&writers.WriterOptions{Metadata: m})
|
||||
err := w.WriteDatabase(db)
|
||||
return w.executionReport, err
|
||||
}
|
||||
|
||||
func columnExists(t *testing.T, conn *pgx.Conn, schema, table, column string) bool {
|
||||
t.Helper()
|
||||
var ok bool
|
||||
err := conn.QueryRow(context.Background(),
|
||||
`SELECT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema=$1 AND table_name=$2 AND column_name=$3)`,
|
||||
schema, table, column).Scan(&ok)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return ok
|
||||
}
|
||||
|
||||
func TestLive_WriteDatabaseEmptyThenIdenticalThenDrifted(t *testing.T) {
|
||||
connString := liveWriterConn(t)
|
||||
schema, conn := liveWriterSchema(t, connString)
|
||||
reportPath := filepath.Join(t.TempDir(), "report.json")
|
||||
meta := map[string]interface{}{"report_path": reportPath}
|
||||
|
||||
// Empty database: schema and table are created.
|
||||
rep, err := runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), meta)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if rep.FailedStatements != 0 || rep.ExecutedStatements == 0 {
|
||||
t.Fatalf("first run report: %+v", rep)
|
||||
}
|
||||
if !columnExists(t, conn, schema, "accounts", "name") {
|
||||
t.Fatal("column name not created")
|
||||
}
|
||||
data, err := os.ReadFile(reportPath)
|
||||
if err != nil {
|
||||
t.Fatalf("report not written: %v", err)
|
||||
}
|
||||
var onDisk ExecutionReport
|
||||
if err := json.Unmarshal(data, &onDisk); err != nil || onDisk.TotalStatements != rep.TotalStatements {
|
||||
t.Errorf("report on disk mismatch: %v %+v", err, onDisk)
|
||||
}
|
||||
created := false
|
||||
for _, s := range rep.Schemas {
|
||||
for _, tb := range s.Tables {
|
||||
if tb.Name == "accounts" && tb.Created {
|
||||
created = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if !created {
|
||||
t.Errorf("table creation not tracked: %+v", rep.Schemas)
|
||||
}
|
||||
|
||||
// Identical database: nothing to execute.
|
||||
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text"}), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if rep.TotalStatements != 0 {
|
||||
t.Errorf("identical DB must produce no statements, got %d", rep.TotalStatements)
|
||||
}
|
||||
|
||||
// Drifted database: only the new column is added.
|
||||
rep, err = runLiveWrite(t, connString, liveModel(schema, map[string]string{"name": "text", "email": "text"}), nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if rep.FailedStatements != 0 || rep.TotalStatements == 0 {
|
||||
t.Errorf("drift report: %+v", rep)
|
||||
}
|
||||
if !columnExists(t, conn, schema, "accounts", "email") {
|
||||
t.Error("drifted column email not added")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_WriteDatabaseFailedStatementContinues(t *testing.T) {
|
||||
connString := liveWriterConn(t)
|
||||
schema, _ := liveWriterSchema(t, connString)
|
||||
reportPath := filepath.Join(t.TempDir(), "report.json")
|
||||
|
||||
db := liveModel(schema, map[string]string{"bad": "no_such_type_xyz"})
|
||||
rep, err := runLiveWrite(t, connString, db, map[string]interface{}{"full_ddl": true, "report_path": reportPath})
|
||||
if err != nil {
|
||||
t.Fatalf("failed statements must not abort the run: %v", err)
|
||||
}
|
||||
if rep.FailedStatements == 0 || len(rep.Errors) != rep.FailedStatements {
|
||||
t.Fatalf("expected recorded failures: %+v", rep)
|
||||
}
|
||||
e := rep.Errors[0]
|
||||
if e.StatementNumber == 0 || e.Statement == "" || e.Error == "" {
|
||||
t.Errorf("incomplete error entry: %+v", e)
|
||||
}
|
||||
if _, err := os.Stat(reportPath); err != nil {
|
||||
t.Errorf("report must be written even on failures: %v", err)
|
||||
}
|
||||
failedTable := false
|
||||
for _, s := range rep.Schemas {
|
||||
for _, tb := range s.Tables {
|
||||
if tb.Name == "accounts" && !tb.Created && tb.Error != "" {
|
||||
failedTable = true
|
||||
}
|
||||
}
|
||||
}
|
||||
if !failedTable {
|
||||
t.Errorf("failed table creation not tracked: %+v", rep.Schemas)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_WriteDatabaseFlattenFallsBackToFullDDL(t *testing.T) {
|
||||
connString := liveWriterConn(t)
|
||||
schema, conn := liveWriterSchema(t, connString)
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{
|
||||
FlattenSchema: true,
|
||||
Metadata: map[string]interface{}{"connection_string": connString},
|
||||
})
|
||||
if err := w.WriteDatabase(liveModel(schema, nil)); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// Flattened output lands in public as <schema>_<table>.
|
||||
flat := "public." + schema + "_accounts"
|
||||
t.Cleanup(func() { _, _ = conn.Exec(context.Background(), "DROP TABLE IF EXISTS "+flat+" CASCADE") })
|
||||
var ok bool
|
||||
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", flat).Scan(&ok); err != nil || !ok {
|
||||
t.Errorf("flattened table %s not created (ok=%v err=%v)", flat, ok, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGenerateLiveDiffStatements_FlattenRejected(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{FlattenSchema: true})
|
||||
if _, err := w.generateLiveDiffStatements(models.InitDatabase("x"), "postgres://unused"); err == nil {
|
||||
t.Error("flatten must be rejected before connecting")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExecuteStatements_ConnectFailure(t *testing.T) {
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
w.executionReport = &ExecutionReport{}
|
||||
err := w.executeStatements([]string{"SELECT 1"}, "postgres://nobody:x@127.0.0.1:1/none?connect_timeout=1")
|
||||
if err == nil {
|
||||
t.Error("expected connect failure")
|
||||
}
|
||||
if w.executionReport.TotalStatements != 1 {
|
||||
t.Errorf("total not recorded: %+v", w.executionReport)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLive_ExecuteStatementsSkipsCommentsAndBlank(t *testing.T) {
|
||||
connString := liveWriterConn(t)
|
||||
schema, conn := liveWriterSchema(t, connString)
|
||||
|
||||
w := NewWriter(&writers.WriterOptions{})
|
||||
w.executionReport = &ExecutionReport{}
|
||||
stmts := []string{
|
||||
"-- Schema: " + schema,
|
||||
" ",
|
||||
"CREATE SCHEMA " + schema,
|
||||
"CREATE TABLE " + schema + ".t (id int)",
|
||||
"-- plain comment",
|
||||
}
|
||||
if err := w.executeStatements(stmts, connString); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
r := w.executionReport
|
||||
if r.ExecutedStatements != 2 || r.FailedStatements != 0 || r.TotalStatements != 5 {
|
||||
t.Errorf("counts: %+v", r)
|
||||
}
|
||||
if len(r.Schemas) != 1 || r.Schemas[0].Name != schema || len(r.Schemas[0].Tables) != 1 || !r.Schemas[0].Tables[0].Created {
|
||||
t.Errorf("schema tracking: %+v", r.Schemas)
|
||||
}
|
||||
var ok bool
|
||||
if err := conn.QueryRow(context.Background(), "SELECT to_regclass($1) IS NOT NULL", schema+".t").Scan(&ok); err != nil || !ok {
|
||||
t.Errorf("table not created: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,288 @@
|
||||
package pgsql
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
||||
)
|
||||
|
||||
func TestExtractTableNameFromCreate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, in, want string
|
||||
}{
|
||||
{"not create table", "SELECT 1", ""},
|
||||
{"plain", "CREATE TABLE users (id int)", "users"},
|
||||
{"qualified", "CREATE TABLE public.users (id int)", "users"},
|
||||
{"if not exists", "CREATE TABLE IF NOT EXISTS public.users (id int)", "users"},
|
||||
{"lowercase", "create table users(id int)", "users"},
|
||||
{"newline", "CREATE TABLE\npublic.t\n(id int)", "t"},
|
||||
{"no name", "CREATE TABLE", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := extractTableNameFromCreate(tt.in); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestTruncateStatement(t *testing.T) {
|
||||
short := strings.Repeat("a", 200)
|
||||
if got := truncateStatement(short); got != short {
|
||||
t.Errorf("200-char statement must not be truncated")
|
||||
}
|
||||
long := strings.Repeat("a", 201)
|
||||
got := truncateStatement(long)
|
||||
if got != strings.Repeat("a", 200)+"..." {
|
||||
t.Errorf("unexpected truncation: len=%d", len(got))
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetCurrentTimestamp(t *testing.T) {
|
||||
ts := getCurrentTimestamp()
|
||||
if len(ts) != len("2006-01-02 15:04:05") || ts[4] != '-' || ts[10] != ' ' || ts[13] != ':' {
|
||||
t.Errorf("unexpected timestamp format %q", ts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractStatementContext(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, in, want string
|
||||
}{
|
||||
{"do block", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.columns WHERE table_schema = 'public' AND table_name = 'users' AND column_name = 'email') THEN NULL; END IF; END $$;`, "public.users (email)"},
|
||||
{"do block constraint", `DO $$ BEGIN IF NOT EXISTS (SELECT 1 FROM information_schema.table_constraints WHERE table_schema = 'public' AND table_name = 'users' AND constraint_name = 'uq_email') THEN NULL; END IF; END $$;`, "public.users [uq_email]"},
|
||||
{"add column", `ALTER TABLE public.users ADD COLUMN "email" text`, "public.users (email)"},
|
||||
{"alter column", `ALTER TABLE users ALTER COLUMN age SET NOT NULL`, "users (age)"},
|
||||
{"add constraint", `ALTER TABLE public.users ADD CONSTRAINT uq_email UNIQUE (email)`, "public.users [uq_email]"},
|
||||
{"drop constraint", `ALTER TABLE public.users DROP CONSTRAINT "uq_email"`, "public.users [uq_email]"},
|
||||
{"alter table plain", `ALTER TABLE public.users RENAME TO people`, "public.users"},
|
||||
{"create table", `CREATE TABLE public.users (id int)`, "public.users"},
|
||||
{"create table if not exists", `CREATE TABLE IF NOT EXISTS "public"."users" (id int)`, "public.users"},
|
||||
{"create schema", `CREATE SCHEMA IF_x;`, "IF_x"},
|
||||
{"create index", `CREATE INDEX idx ON public.users (email)`, "public.users"},
|
||||
{"create unique index", `CREATE UNIQUE INDEX idx ON users (email)`, "users"},
|
||||
{"create index without on", `CREATE INDEX idx`, ""},
|
||||
{"comment on table", `COMMENT ON TABLE public.users IS 'x'`, "public.users"},
|
||||
{"comment on column", `COMMENT ON COLUMN public.users.email IS 'x'`, "public.users.email"},
|
||||
{"unknown", `DROP TABLE users`, ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := extractStatementContext(tt.in); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestExtractSQLStringValue(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, stmt, key, want string
|
||||
}{
|
||||
{"basic", "WHERE table_name = 'users'", "table_name", "users"},
|
||||
{"case-insensitive key", "WHERE TABLE_NAME='users'", "table_name", "users"},
|
||||
{"missing key", "WHERE a = 'b'", "table_name", ""},
|
||||
{"no equals", "table_name is 'x'", "table_name", ""},
|
||||
{"equals too far", "table_name abcdefgh = 'x'", "table_name", ""},
|
||||
{"not quoted", "table_name = users", "table_name", ""},
|
||||
{"unterminated", "table_name = 'users", "table_name", ""},
|
||||
{"empty after key", "table_name", "table_name", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := extractSQLStringValue(tt.stmt, tt.key); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseQualifiedIdent(t *testing.T) {
|
||||
tests := []struct {
|
||||
in, schema, name string
|
||||
}{
|
||||
{"users (id int)", "", "users"},
|
||||
{"public.users (id int)", "public", "users"},
|
||||
{`"public"."users" (id int)`, "public", "users"},
|
||||
{`"users"`, "", "users"},
|
||||
{"", "", ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
s, n := parseQualifiedIdent(tt.in)
|
||||
if s != tt.schema || n != tt.name {
|
||||
t.Errorf("parseQualifiedIdent(%q) = (%q,%q), want (%q,%q)", tt.in, s, n, tt.schema, tt.name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestFirstBareIdentAndHelpers(t *testing.T) {
|
||||
bare := map[string]string{
|
||||
"": "",
|
||||
" ": "",
|
||||
"abc": "abc",
|
||||
"abc def": "abc",
|
||||
"abc(def)": "abc",
|
||||
"abc,def": "abc",
|
||||
"abc;": "abc",
|
||||
"\n abc\tdef": "abc",
|
||||
`"a b" c`: `"a`,
|
||||
" tbl (x int)": "tbl",
|
||||
}
|
||||
for in, want := range bare {
|
||||
if got := firstBareIdent(in); got != want {
|
||||
t.Errorf("firstBareIdent(%q) = %q, want %q", in, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
if got := stripQuotes(`"abc"`); got != "abc" {
|
||||
t.Errorf("stripQuotes = %q", got)
|
||||
}
|
||||
if got := stripQuotes("abc"); got != "abc" {
|
||||
t.Errorf("stripQuotes unquoted = %q", got)
|
||||
}
|
||||
|
||||
stmt := `ALTER TABLE t add column "c1" text`
|
||||
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "ADD COLUMN"); got != "c1" {
|
||||
t.Errorf("firstIdentAfterKeyword = %q", got)
|
||||
}
|
||||
if got := firstIdentAfterKeyword(stmt, strings.ToUpper(stmt), "DROP COLUMN"); got != "" {
|
||||
t.Errorf("missing keyword must return empty, got %q", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBuildStmtContext(t *testing.T) {
|
||||
tests := []struct {
|
||||
schema, table, column, constraint, want string
|
||||
}{
|
||||
{"", "", "", "", ""},
|
||||
{"s", "t", "", "", "s.t"},
|
||||
{"", "t", "", "", "t"},
|
||||
{"s", "", "", "", ""},
|
||||
{"s", "t", "c", "", "s.t (c)"},
|
||||
{"s", "t", "", "k", "s.t [k]"},
|
||||
{"s", "t", "c", "k", "s.t (c) [k]"},
|
||||
{"", "", "c", "", "(c)"},
|
||||
{"", "", "", "k", "[k]"},
|
||||
{"", "", "c", "k", "(c) [k]"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := buildStmtContext(tt.schema, tt.table, tt.column, tt.constraint); got != tt.want {
|
||||
t.Errorf("buildStmtContext(%q,%q,%q,%q) = %q, want %q", tt.schema, tt.table, tt.column, tt.constraint, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestDetectStatementType(t *testing.T) {
|
||||
tests := []struct {
|
||||
name, in, want string
|
||||
}{
|
||||
{"do unique", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT u UNIQUE (a); END $$", "ADD UNIQUE CONSTRAINT"},
|
||||
{"do fk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id); END $$", "ADD FOREIGN KEY"},
|
||||
{"do pk", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a); END $$", "ADD PRIMARY KEY"},
|
||||
{"do check", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c CHECK (a > 0); END $$", "ADD CHECK CONSTRAINT"},
|
||||
{"do constraint", "DO $$ BEGIN ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a); END $$", "ADD CONSTRAINT"},
|
||||
{"do add column", "DO $$ BEGIN ALTER TABLE t ADD COLUMN c int; END $$", "ADD COLUMN"},
|
||||
{"do drop constraint", "DO $$ BEGIN DROP CONSTRAINT x; END $$", "DROP CONSTRAINT"},
|
||||
{"do other", "DO $$ BEGIN NULL; END $$", "DO BLOCK"},
|
||||
{"create schema", "create schema s", "CREATE SCHEMA"},
|
||||
{"create sequence", "CREATE SEQUENCE s", "CREATE SEQUENCE"},
|
||||
{"create table", "CREATE TABLE t ()", "CREATE TABLE"},
|
||||
{"create index", "CREATE INDEX i ON t(a)", "CREATE INDEX"},
|
||||
{"create unique index", "CREATE UNIQUE INDEX i ON t(a)", "CREATE UNIQUE INDEX"},
|
||||
{"alter fk", "ALTER TABLE t ADD CONSTRAINT f FOREIGN KEY (a) REFERENCES x(id)", "ADD FOREIGN KEY"},
|
||||
{"alter pk", "ALTER TABLE t ADD CONSTRAINT p PRIMARY KEY (a)", "ADD PRIMARY KEY"},
|
||||
{"alter unique", "ALTER TABLE t ADD CONSTRAINT u UNIQUE (a)", "ADD UNIQUE CONSTRAINT"},
|
||||
{"alter check", "ALTER TABLE t ADD CONSTRAINT c CHECK (a>0)", "ADD CHECK CONSTRAINT"},
|
||||
{"alter constraint", "ALTER TABLE t ADD CONSTRAINT c EXCLUDE (a)", "ADD CONSTRAINT"},
|
||||
{"alter add column", "ALTER TABLE t ADD COLUMN c int", "ADD COLUMN"},
|
||||
{"alter drop constraint", "ALTER TABLE t DROP CONSTRAINT c", "DROP CONSTRAINT"},
|
||||
{"alter column", "ALTER TABLE t ALTER COLUMN c TYPE int", "ALTER COLUMN"},
|
||||
{"alter table", "ALTER TABLE t RENAME TO u", "ALTER TABLE"},
|
||||
{"comment table", "COMMENT ON TABLE t IS 'x'", "COMMENT ON TABLE"},
|
||||
{"comment column", "COMMENT ON COLUMN t.c IS 'x'", "COMMENT ON COLUMN"},
|
||||
{"drop table", "DROP TABLE t", "DROP TABLE"},
|
||||
{"drop index", "DROP INDEX i", "DROP INDEX"},
|
||||
{"default", "SELECT 1", "SQL"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := detectStatementType(tt.in); got != tt.want {
|
||||
t.Errorf("got %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteAndFinishReport(t *testing.T) {
|
||||
report := &ExecutionReport{
|
||||
TotalStatements: 3,
|
||||
ExecutedStatements: 2,
|
||||
FailedStatements: 1,
|
||||
Schemas: []SchemaReport{{Name: "public", Tables: []TableReport{
|
||||
{Name: "a", Created: true},
|
||||
{Name: "b", Created: false, Error: "boom"},
|
||||
}}},
|
||||
Errors: []ExecutionError{{StatementNumber: 3, Statement: "CREATE TABLE b ()", Error: "boom"}},
|
||||
StartTime: "s",
|
||||
EndTime: "e",
|
||||
}
|
||||
|
||||
path := filepath.Join(t.TempDir(), "report.json")
|
||||
w := &Writer{
|
||||
options: &writers.WriterOptions{Metadata: map[string]interface{}{"report_path": path}},
|
||||
executionReport: report,
|
||||
}
|
||||
if err := w.finishReport(); err != nil {
|
||||
t.Fatalf("finishReport: %v", err)
|
||||
}
|
||||
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
t.Fatalf("report not written: %v", err)
|
||||
}
|
||||
var got ExecutionReport
|
||||
if err := json.Unmarshal(data, &got); err != nil {
|
||||
t.Fatalf("invalid report JSON: %v", err)
|
||||
}
|
||||
if got.TotalStatements != 3 || got.FailedStatements != 1 || len(got.Errors) != 1 ||
|
||||
len(got.Schemas) != 1 || len(got.Schemas[0].Tables) != 2 || got.Schemas[0].Tables[1].Error != "boom" {
|
||||
t.Errorf("report round-trip mismatch: %+v", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestFinishReportNoPathAndSuccess(t *testing.T) {
|
||||
w := &Writer{
|
||||
options: &writers.WriterOptions{},
|
||||
executionReport: &ExecutionReport{TotalStatements: 1, ExecutedStatements: 1},
|
||||
}
|
||||
if err := w.finishReport(); err != nil {
|
||||
t.Errorf("finishReport without path: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWriteReportBadPath(t *testing.T) {
|
||||
w := &Writer{options: &writers.WriterOptions{}, executionReport: &ExecutionReport{}}
|
||||
if err := w.writeReport(filepath.Join(t.TempDir(), "missing", "r.json")); err == nil {
|
||||
t.Error("expected error for unwritable path")
|
||||
}
|
||||
// finishReport must swallow the report error.
|
||||
w.options.Metadata = map[string]interface{}{"report_path": filepath.Join(t.TempDir(), "missing", "r.json")}
|
||||
if err := w.finishReport(); err != nil {
|
||||
t.Errorf("finishReport must not fail on report write error: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTemplateFilterAndMapFuncPassthrough(t *testing.T) {
|
||||
in := []string{"a", "b"}
|
||||
if got := filter(in, "X").([]string); len(got) != 2 {
|
||||
t.Errorf("filter must return slice unchanged")
|
||||
}
|
||||
if got := mapFunc("v", "upper"); got != "v" {
|
||||
t.Errorf("mapFunc must return value unchanged, got %v", got)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user