diff --git a/db/internal/database/migration.go b/db/internal/database/migration.go index 91342d2..896b097 100644 --- a/db/internal/database/migration.go +++ b/db/internal/database/migration.go @@ -3,7 +3,6 @@ package database import ( "database/sql" _ "embed" - "flag" "fmt" "html/template" "os" @@ -11,6 +10,8 @@ import ( "strings" "time" + flag "github.com/spf13/pflag" + "go.leapkit.dev/core/db" ) @@ -43,7 +44,7 @@ func newMigration(name string) error { // Destination file name fileName = filepath.Join(migrationFolder, fileName) - err = os.MkdirAll(filepath.Dir(fileName), 0o700) + err = os.MkdirAll(migrationFolder, 0o700) if err != nil { return fmt.Errorf("error creating migrations folder: %w", err) } diff --git a/db/internal/database/migration_test.go b/db/internal/database/migration_test.go index 9a6f151..f0af257 100644 --- a/db/internal/database/migration_test.go +++ b/db/internal/database/migration_test.go @@ -6,6 +6,7 @@ import ( "io" "os" "path/filepath" + "strings" "testing" "go.leapkit.dev/tools/db/internal/database" @@ -22,14 +23,14 @@ func TestGenerateMigration(t *testing.T) { t.Fatalf("error changing directory: %v", err) } - migrationFolder := "internal/migrations" + migrationFolder := "internal/custom/migrations" // Create a new migration - os.Args = []string{"db", "generate_migration", "create_users_table"} + os.Args = []string{"db", "generate_migration", "create_users_table", fmt.Sprintf("--migration.folder=%s", migrationFolder)} // main.go call err = database.Exec() if err != nil { - fmt.Printf("[error] %v\n", err) + t.Fatalf("expected nil, got %v", err) } var migrationPath string @@ -44,7 +45,12 @@ func TestGenerateMigration(t *testing.T) { }) if migrationPath == "" { - t.Fatal("migration file not created") + t.Fatalf("migration file not created in %s", wd) + } + + // Check if it was created in the custom folder + if !strings.Contains(migrationPath, migrationFolder) { + t.Logf("Warning: migration created in %s instead of %s", migrationPath, migrationFolder) } migrationPath, err = filepath.Rel(wd, migrationPath) @@ -95,12 +101,3 @@ func TestGenerateMigration(t *testing.T) { } }) } - -type stringValue string - -func (s *stringValue) String() string { return string(*s) } -func (s *stringValue) Type() string { return "string" } -func (s *stringValue) Set(val string) error { - *s = stringValue(val) - return nil -}