309 lines
9.7 KiB
Go
309 lines
9.7 KiB
Go
package drizzle
|
|
|
|
import (
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"testing"
|
|
|
|
"git.warky.dev/wdevs/relspecgo/pkg/models"
|
|
"git.warky.dev/wdevs/relspecgo/pkg/readers"
|
|
drizzlereader "git.warky.dev/wdevs/relspecgo/pkg/readers/drizzle"
|
|
"git.warky.dev/wdevs/relspecgo/pkg/writers"
|
|
)
|
|
|
|
const drizzleFixture = "../../../tests/assets/drizzle/schema.ts"
|
|
|
|
func fixtureDB(t *testing.T) *models.Database {
|
|
t.Helper()
|
|
db, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: drizzleFixture}).ReadDatabase()
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return db
|
|
}
|
|
|
|
// shopDB builds a database with an enum, FK, unique, index and varied defaults.
|
|
func shopDB() *models.Database {
|
|
s := models.InitSchema("public")
|
|
s.Enums = append(s.Enums, &models.Enum{Name: "role", Schema: "public", Values: []string{"admin", "user"}})
|
|
|
|
users := models.InitTable("users", "public")
|
|
id := models.InitColumn("id", "users", "public")
|
|
id.Type, id.IsPrimaryKey, id.NotNull, id.AutoIncrement = "integer", true, true, true
|
|
email := models.InitColumn("email", "users", "public")
|
|
email.Type, email.NotNull = "varchar(255)", true
|
|
role := models.InitColumn("role", "users", "public")
|
|
role.Type, role.NotNull, role.Default = "role", true, "user"
|
|
active := models.InitColumn("active", "users", "public")
|
|
active.Type, active.Default = "boolean", true
|
|
created := models.InitColumn("created_at", "users", "public")
|
|
created.Type, created.Default = "timestamp", "now()"
|
|
score := models.InitColumn("score", "users", "public")
|
|
score.Type, score.Default = "integer", "10"
|
|
for _, c := range []*models.Column{id, email, role, active, created, score} {
|
|
users.Columns[c.Name] = c
|
|
}
|
|
uq := models.InitConstraint("uq_email", models.UniqueConstraint)
|
|
uq.Columns = []string{"email"}
|
|
users.Constraints["uq_email"] = uq
|
|
ix := models.InitIndex("idx_role", "users", "public")
|
|
ix.Columns = []string{"role"}
|
|
users.Indexes["idx_role"] = ix
|
|
|
|
posts := models.InitTable("blog_posts", "public")
|
|
pid := models.InitColumn("id", "blog_posts", "public")
|
|
pid.Type, pid.IsPrimaryKey, pid.NotNull = "uuid", true, true
|
|
pid.Default = "gen_random_uuid()"
|
|
author := models.InitColumn("author_id", "blog_posts", "public")
|
|
author.Type, author.NotNull = "integer", true
|
|
posts.Columns["id"], posts.Columns["author_id"] = pid, author
|
|
fk := models.InitConstraint("fk_author", models.ForeignKeyConstraint)
|
|
fk.Columns, fk.ReferencedTable, fk.ReferencedSchema, fk.ReferencedColumns = []string{"author_id"}, "users", "public", []string{"id"}
|
|
posts.Constraints["fk_author"] = fk
|
|
|
|
s.Tables = append(s.Tables, users, posts)
|
|
db := models.InitDatabase("shop")
|
|
db.Schemas = append(db.Schemas, s)
|
|
return db
|
|
}
|
|
|
|
func writeSingle(t *testing.T, db *models.Database) string {
|
|
t.Helper()
|
|
out := filepath.Join(t.TempDir(), "schema.ts")
|
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(db); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
b, err := os.ReadFile(out)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
return string(b)
|
|
}
|
|
|
|
func TestSingleFile_Content(t *testing.T) {
|
|
got := writeSingle(t, shopDB())
|
|
for _, want := range []string{
|
|
"drizzle-orm/pg-core", "pgEnum('role'", "'admin'", "'user'",
|
|
"pgTable('users'", "pgTable('blog_posts'", "blogPosts",
|
|
".primaryKey()", ".notNull()", ".unique()", ".references(() => users.id)",
|
|
"default(true)", "default(10)", "sql`now()`", "sql`gen_random_uuid()`",
|
|
"varchar", "uuid(",
|
|
} {
|
|
if !strings.Contains(got, want) {
|
|
t.Errorf("output missing %q\n%s", want, got)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestSingleFile_Deterministic(t *testing.T) {
|
|
first := writeSingle(t, shopDB())
|
|
for i := 0; i < 15; i++ {
|
|
if got := writeSingle(t, shopDB()); got != first {
|
|
t.Fatalf("output differs on run %d", i)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestFixtureRoundTrip(t *testing.T) {
|
|
db := fixtureDB(t)
|
|
got := writeSingle(t, db)
|
|
if got == "" {
|
|
t.Fatal("empty output")
|
|
}
|
|
out := filepath.Join(t.TempDir(), "again.ts")
|
|
if err := os.WriteFile(out, []byte(got), 0o644); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
again, err := drizzlereader.NewReader(&readers.ReaderOptions{FilePath: out}).ReadDatabase()
|
|
if err != nil {
|
|
t.Fatalf("re-read: %v", err)
|
|
}
|
|
count := func(d *models.Database) (n int) {
|
|
for _, s := range d.Schemas {
|
|
n += len(s.Tables)
|
|
}
|
|
return
|
|
}
|
|
if count(db) != count(again) {
|
|
t.Errorf("tables: %d -> %d", count(db), count(again))
|
|
}
|
|
}
|
|
|
|
func TestMultiFile(t *testing.T) {
|
|
dir := filepath.Join(t.TempDir(), "schema")
|
|
w := NewWriter(&writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": true}})
|
|
if err := w.WriteDatabase(shopDB()); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
for _, f := range []string{"enums.ts", "users.ts", "blog_posts.ts"} {
|
|
if _, err := os.Stat(filepath.Join(dir, f)); err != nil {
|
|
t.Errorf("%s not written: %v", f, err)
|
|
}
|
|
}
|
|
users, _ := os.ReadFile(filepath.Join(dir, "users.ts"))
|
|
if !strings.Contains(string(users), "from './enums'") {
|
|
t.Errorf("users.ts must import its enum:\n%s", users)
|
|
}
|
|
posts, _ := os.ReadFile(filepath.Join(dir, "blog_posts.ts"))
|
|
if strings.Contains(string(posts), "from './enums'") {
|
|
t.Errorf("blog_posts.ts uses no enum:\n%s", posts)
|
|
}
|
|
enums, _ := os.ReadFile(filepath.Join(dir, "enums.ts"))
|
|
if !strings.Contains(string(enums), "pgEnum('role'") {
|
|
t.Errorf("enums.ts:\n%s", enums)
|
|
}
|
|
}
|
|
|
|
func TestMultiFile_RequiresOutputPath(t *testing.T) {
|
|
w := NewWriter(&writers.WriterOptions{Metadata: map[string]any{"multi_file": true}})
|
|
if err := w.WriteDatabase(shopDB()); err == nil || !strings.Contains(err.Error(), "output path is required") {
|
|
t.Errorf("got %v", err)
|
|
}
|
|
}
|
|
|
|
func TestShouldUseMultiFile(t *testing.T) {
|
|
dir := t.TempDir()
|
|
tests := []struct {
|
|
name string
|
|
opts writers.WriterOptions
|
|
want bool
|
|
}{
|
|
{"stdout", writers.WriterOptions{}, false},
|
|
{"explicit true", writers.WriterOptions{OutputPath: "x.ts", Metadata: map[string]any{"multi_file": true}}, true},
|
|
{"explicit false", writers.WriterOptions{OutputPath: dir, Metadata: map[string]any{"multi_file": false}}, false},
|
|
{"ts file", writers.WriterOptions{OutputPath: "schema.ts"}, false},
|
|
{"trailing slash", writers.WriterOptions{OutputPath: "out/"}, true},
|
|
{"trailing backslash", writers.WriterOptions{OutputPath: `out\`}, true},
|
|
{"existing dir", writers.WriterOptions{OutputPath: dir}, true},
|
|
{"nonexistent no ext", writers.WriterOptions{OutputPath: filepath.Join(dir, "nope")}, false},
|
|
}
|
|
for _, tt := range tests {
|
|
t.Run(tt.name, func(t *testing.T) {
|
|
opts := tt.opts
|
|
if got := NewWriter(&opts).shouldUseMultiFile(); got != tt.want {
|
|
t.Errorf("got %v, want %v", got, tt.want)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
func TestWriteSchemaAndTable(t *testing.T) {
|
|
db := shopDB()
|
|
dir := t.TempDir()
|
|
|
|
sOut := filepath.Join(dir, "s.ts")
|
|
if err := NewWriter(&writers.WriterOptions{OutputPath: sOut}).WriteSchema(db.Schemas[0]); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
tOut := filepath.Join(dir, "t.ts")
|
|
if err := NewWriter(&writers.WriterOptions{OutputPath: tOut}).WriteTable(db.Schemas[0].Tables[0]); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
b, _ := os.ReadFile(tOut)
|
|
if !strings.Contains(string(b), "pgTable('users'") {
|
|
t.Errorf("table output:\n%s", b)
|
|
}
|
|
}
|
|
|
|
func TestWriteDatabase_BadOutputPath(t *testing.T) {
|
|
out := filepath.Join(t.TempDir(), "missing", "x.ts")
|
|
if err := NewWriter(&writers.WriterOptions{OutputPath: out}).WriteDatabase(shopDB()); err == nil {
|
|
t.Error("expected error")
|
|
}
|
|
}
|
|
|
|
func TestFormatDefaultValue(t *testing.T) {
|
|
tm := NewTypeMapper()
|
|
tests := []struct {
|
|
in any
|
|
want string
|
|
}{
|
|
{"now()", "sql`now()`"},
|
|
{"CURRENT_TIMESTAMP", "sql`now()`"},
|
|
{"gen_random_uuid()", "sql`gen_random_uuid()`"},
|
|
{"uuid_generate_v4()", "sql`gen_random_uuid()`"},
|
|
{"42", "42"},
|
|
{"-1.5", "-1.5"},
|
|
{"it's", `'it\'s'`},
|
|
{"plain", "'plain'"},
|
|
{true, "true"},
|
|
{false, "false"},
|
|
{7, "7"},
|
|
{int64(8), "8"},
|
|
{2.5, "2.5"},
|
|
}
|
|
for _, tt := range tests {
|
|
if got := tm.formatDefaultValue(tt.in); got != tt.want {
|
|
t.Errorf("formatDefaultValue(%#v) = %q, want %q", tt.in, got, tt.want)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestIsNumericString(t *testing.T) {
|
|
for in, want := range map[string]bool{"": false, "1": true, "-1": true, "1.5": true, "1a": false, "a": false, "1-": false} {
|
|
if got := isNumericString(in); got != want {
|
|
t.Errorf("isNumericString(%q) = %v", in, got)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestBuildReferencesChain(t *testing.T) {
|
|
tm := NewTypeMapper()
|
|
fk := &models.Constraint{ReferencedColumns: []string{"id"}}
|
|
if got := tm.BuildReferencesChain(fk, "blog_posts"); got != "references(() => blogPosts.id)" {
|
|
t.Errorf("got %q", got)
|
|
}
|
|
if got := tm.BuildReferencesChain(&models.Constraint{}, "x"); got != "" {
|
|
t.Errorf("no columns: %q", got)
|
|
}
|
|
}
|
|
|
|
func TestSortHelpers(t *testing.T) {
|
|
idxs := map[string]*models.Index{
|
|
"b": {Name: "b"}, "a": {Name: "a"}, "s2": {Name: "z", Sequence: 2}, "s1": {Name: "y", Sequence: 1},
|
|
}
|
|
got := sortIndexes(idxs)
|
|
if len(got) != 4 {
|
|
t.Fatalf("len %d", len(got))
|
|
}
|
|
// Items with a sequence are ordered by it relative to each other.
|
|
pos := map[string]int{}
|
|
for i, ix := range got {
|
|
pos[ix.Name] = i
|
|
}
|
|
if pos["y"] > pos["z"] {
|
|
t.Errorf("sequence order violated: %v", pos)
|
|
}
|
|
|
|
cons := sortConstraints(map[string]*models.Constraint{"b": {Name: "b"}, "a": {Name: "a"}})
|
|
if len(cons) != 2 || cons[0].Name != "a" {
|
|
t.Errorf("sortConstraints: %+v", cons)
|
|
}
|
|
|
|
strs := []string{"c", "a", "b"}
|
|
sortStrings(strs)
|
|
if strings.Join(strs, "") != "abc" {
|
|
t.Errorf("sortStrings: %v", strs)
|
|
}
|
|
}
|
|
|
|
func TestEnumColumnCallsConstant(t *testing.T) {
|
|
out := filepath.Join(t.TempDir(), "schema.ts")
|
|
w := NewWriter(&writers.WriterOptions{OutputPath: out})
|
|
if err := w.WriteDatabase(&models.Database{Name: "d", Schemas: []*models.Schema{shopDB().Schemas[0]}}); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
b, err := os.ReadFile(out)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
got := string(b)
|
|
if !strings.Contains(got, "role('role')") {
|
|
t.Errorf("enum column should call constant:\n%s", got)
|
|
}
|
|
if strings.Contains(got, "pgEnum('role')(") {
|
|
t.Errorf("invalid pgEnum(...)(...) syntax emitted:\n%s", got)
|
|
}
|
|
}
|