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) } } func TestSplitSQLSemicolonInStringLiteral(t *testing.T) { sql := `SELECT 'hello; world'; INSERT INTO t VALUES (1);` stmts := nonEmpty(splitSQL(sql)) if len(stmts) != 2 { t.Fatalf("got %d statements, want 2: %#v", len(stmts), stmts) } } func TestSplitSQLDollarSignInStringLiteral(t *testing.T) { sql := `SELECT '$100'; SELECT 2;` stmts := nonEmpty(splitSQL(sql)) if len(stmts) != 2 { t.Fatalf("got %d statements, want 2: %#v", len(stmts), stmts) } } func TestSplitSQLBlockComment(t *testing.T) { sql := `SELECT 1; /* block; with; semicolons */ SELECT 2;` stmts := nonEmpty(splitSQL(sql)) if len(stmts) != 2 { t.Fatalf("got %d statements, want 2: %#v", len(stmts), stmts) } } func TestSplitSQLBlockCommentWithDollarQuote(t *testing.T) { sql := `/* $$ not a dollar quote */ SELECT 1;` stmts := nonEmpty(splitSQL(sql)) if len(stmts) != 1 { t.Fatalf("got %d statements, want 1: %#v", len(stmts), stmts) } } func TestSplitSQLDoubledQuoteInString(t *testing.T) { sql := `SELECT 'O''Brien'; SELECT 2;` stmts := nonEmpty(splitSQL(sql)) if len(stmts) != 2 { t.Fatalf("got %d statements, want 2: %#v", len(stmts), stmts) } } func TestSplitSQLEmptyInput(t *testing.T) { stmts := nonEmpty(splitSQL("")) if len(stmts) != 0 { t.Fatalf("got %d statements, want 0", len(stmts)) } } func TestSplitSQLNoSemicolon(t *testing.T) { stmts := nonEmpty(splitSQL("SELECT 1")) if len(stmts) != 1 { t.Fatalf("got %d statements, want 1", len(stmts)) } }