package db import ( "strings" "testing" ) func nonEmpty(stmts []string) []string { var out []string for _, s := range stmts { if strings.TrimSpace(s) != "" { out = append(out, s) } } return out } func TestSplitSQLBasic(t *testing.T) { stmts := nonEmpty(splitSQL("CREATE TABLE a (id int); CREATE TABLE b (id int);")) if len(stmts) != 2 { t.Fatalf("got %d statements, want 2: %#v", len(stmts), stmts) } } func TestSplitSQLDollarQuotedFunction(t *testing.T) { sql := `CREATE FUNCTION f() RETURNS int AS $$ SELECT 1; SELECT 2; $$ LANGUAGE sql; CREATE TABLE t (id int);` stmts := nonEmpty(splitSQL(sql)) if len(stmts) != 2 { t.Fatalf("got %d statements, want 2: %#v", len(stmts), stmts) } if !strings.Contains(stmts[0], "SELECT 1; SELECT 2;") { t.Errorf("dollar-quoted body was split: %q", stmts[0]) } } func TestSplitSQLTaggedDollarQuote(t *testing.T) { sql := `DO $body$ BEGIN PERFORM 1; END $body$;SELECT 1;` stmts := nonEmpty(splitSQL(sql)) if len(stmts) != 2 { t.Fatalf("got %d statements, want 2: %#v", len(stmts), stmts) } } func TestSplitSQLSemicolonInComment(t *testing.T) { sql := "-- comment with ; semicolon\nCREATE TABLE t (id int); -- trailing; note\nSELECT 1;" stmts := nonEmpty(splitSQL(sql)) if len(stmts) != 2 { t.Fatalf("got %d statements, want 2: %#v", len(stmts), stmts) } }