Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 3 additions & 2 deletions db/internal/database/migration.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,14 +3,15 @@ package database
import (
"database/sql"
_ "embed"
"flag"
"fmt"
"html/template"
"os"
"path/filepath"
"strings"
"time"

flag "github.com/spf13/pflag"

"go.leapkit.dev/core/db"
)

Expand Down Expand Up @@ -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)
}
Expand Down
23 changes: 10 additions & 13 deletions db/internal/database/migration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"io"
"os"
"path/filepath"
"strings"
"testing"

"go.leapkit.dev/tools/db/internal/database"
Expand All @@ -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
Expand All @@ -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)
Expand Down Expand Up @@ -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
}
Loading