jobs-monorepo/internal/infrastructure/db.go
2026-01-26 17:33:48 +05:00

183 lines
4.5 KiB
Go

package infrastructure
import (
"database/sql"
"errors"
"fmt"
"os"
"path/filepath"
"runtime"
"github.com/golang-migrate/migrate/v4"
"github.com/golang-migrate/migrate/v4/database/postgres"
_ "github.com/golang-migrate/migrate/v4/source/file"
_ "github.com/lib/pq"
)
// Config holds database configuration
type Config struct {
Host string
Port string
User string
Password string
DBName string
SSLMode string
}
// NewConnection creates a new database connection
func NewConnection(config Config) (*sql.DB, error) {
dsn := fmt.Sprintf("host=%s port=%s user=%s password=%s dbname=%s sslmode=%s",
config.Host, config.Port, config.User, config.Password, config.DBName, config.SSLMode)
db, err := sql.Open("postgres", dsn)
if err != nil {
return nil, fmt.Errorf("failed to open database connection: %w", err)
}
if err := db.Ping(); err != nil {
return nil, fmt.Errorf("failed to ping database: %w", err)
}
return db, nil
}
// LoadConfigFromEnv loads database configuration from environment variables
func LoadConfigFromEnv() Config {
return Config{
Host: getEnv("DB_HOST", "localhost"),
Port: getEnv("DB_PORT", "5432"),
User: getEnv("DB_USER", "postgres"),
Password: getEnv("DB_PASSWORD", "password"),
DBName: getEnv("DB_NAME", "linkedin_jobs"),
SSLMode: getEnv("DB_SSLMODE", "disable"),
}
}
func getEnv(key, defaultValue string) string {
if value := os.Getenv(key); value != "" {
return value
}
return defaultValue
}
// RunMigrations runs all pending migrations
func RunMigrations(db *sql.DB) error {
driver, err := postgres.WithInstance(db, &postgres.Config{})
if err != nil {
return fmt.Errorf("failed to create postgres driver: %w", err)
}
// Get the migrations directory path
migrationsPath, err := getMigrationsPath()
if err != nil {
return fmt.Errorf("failed to get migrations path: %w", err)
}
m, err := migrate.NewWithDatabaseInstance(
fmt.Sprintf("file://%s", migrationsPath),
"postgres",
driver,
)
if err != nil {
return fmt.Errorf("failed to create migrate instance: %w", err)
}
if err := m.Up(); err != nil && !errors.Is(err, migrate.ErrNoChange) {
return fmt.Errorf("failed to run migrations: %w", err)
}
return nil
}
// RollbackMigrations rolls back migrations by the specified number of steps
func RollbackMigrations(db *sql.DB, steps int) error {
driver, err := postgres.WithInstance(db, &postgres.Config{})
if err != nil {
return fmt.Errorf("failed to create postgres driver: %w", err)
}
migrationsPath, err := getMigrationsPath()
if err != nil {
return fmt.Errorf("failed to get migrations path: %w", err)
}
m, err := migrate.NewWithDatabaseInstance(
fmt.Sprintf("file://%s", migrationsPath),
"postgres",
driver,
)
if err != nil {
return fmt.Errorf("failed to create migrate instance: %w", err)
}
defer m.Close()
if err := m.Steps(-steps); err != nil && !errors.Is(err, migrate.ErrNoChange) {
return fmt.Errorf("failed to rollback migrations: %w", err)
}
return nil
}
// GetMigrationVersion returns the current migration version
func GetMigrationVersion(db *sql.DB) (uint, bool, error) {
driver, err := postgres.WithInstance(db, &postgres.Config{})
if err != nil {
return 0, false, fmt.Errorf("failed to create postgres driver: %w", err)
}
migrationsPath, err := getMigrationsPath()
if err != nil {
return 0, false, fmt.Errorf("failed to get migrations path: %w", err)
}
m, err := migrate.NewWithDatabaseInstance(
fmt.Sprintf("file://%s", migrationsPath),
"postgres",
driver,
)
if err != nil {
return 0, false, fmt.Errorf("failed to create migrate instance: %w", err)
}
defer m.Close()
version, dirty, err := m.Version()
if err != nil {
if errors.Is(err, migrate.ErrNilVersion) {
return 0, false, nil
}
return 0, false, fmt.Errorf("failed to get migration version: %w", err)
}
return version, dirty, nil
}
func getMigrationsPath() (string, error) {
if envPath := os.Getenv("MIGRATIONS_PATH"); envPath != "" {
if _, err := os.Stat(envPath); err == nil {
return envPath, nil
}
}
paths := []string{
"./internal/migrations",
"../internal/migrations",
"../../internal/migrations",
}
for _, p := range paths {
if _, err := os.Stat(p); err == nil {
return p, nil
}
}
_, filename, _, ok := runtime.Caller(0)
if ok {
internalDir := filepath.Dir(filepath.Dir(filepath.Dir(filename)))
migrationsPath := filepath.Join(internalDir, "migrations")
if _, err := os.Stat(migrationsPath); err == nil {
return migrationsPath, nil
}
}
return "", fmt.Errorf("migrations directory not found")
}