81 lines
1.9 KiB
Go
81 lines
1.9 KiB
Go
package db
|
|
|
|
import (
|
|
"context"
|
|
"database/sql"
|
|
"fmt"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
|
|
_ "github.com/jackc/pgx/v5/stdlib"
|
|
_ "modernc.org/sqlite"
|
|
)
|
|
|
|
func New(ctx context.Context, driver, dsn string) (*sql.DB, error) {
|
|
driverName, err := normalizeDriver(driver)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if driverName == "sqlite" {
|
|
if err := ensureSQLiteDirectory(dsn); err != nil {
|
|
return nil, fmt.Errorf("prepare sqlite path: %w", err)
|
|
}
|
|
}
|
|
|
|
database, err := sql.Open(driverName, dsn)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("open %s database: %w", driverName, err)
|
|
}
|
|
|
|
if driverName == "sqlite" {
|
|
// In-memory databases are connection-local, so keep SQLite single-connection
|
|
// and enable foreign keys for every operation through this handle.
|
|
database.SetMaxOpenConns(1)
|
|
database.SetMaxIdleConns(1)
|
|
if _, err := database.ExecContext(ctx, "PRAGMA foreign_keys = ON"); err != nil {
|
|
database.Close()
|
|
return nil, fmt.Errorf("configure sqlite: %w", err)
|
|
}
|
|
} else {
|
|
database.SetMaxOpenConns(10)
|
|
database.SetMaxIdleConns(1)
|
|
}
|
|
|
|
if err := database.PingContext(ctx); err != nil {
|
|
database.Close()
|
|
return nil, fmt.Errorf("ping %s database: %w", driverName, err)
|
|
}
|
|
|
|
return database, nil
|
|
}
|
|
|
|
func normalizeDriver(driver string) (string, error) {
|
|
switch strings.ToLower(strings.TrimSpace(driver)) {
|
|
case "postgres", "postgresql", "pgx":
|
|
return "pgx", nil
|
|
case "sqlite", "sqlite3":
|
|
return "sqlite", nil
|
|
default:
|
|
return "", fmt.Errorf("unsupported database driver %q (use postgres or sqlite)", driver)
|
|
}
|
|
}
|
|
|
|
func ensureSQLiteDirectory(dsn string) error {
|
|
if dsn == ":memory:" || strings.HasPrefix(dsn, "file::memory:") {
|
|
return nil
|
|
}
|
|
|
|
path := strings.SplitN(dsn, "?", 2)[0]
|
|
path = strings.TrimPrefix(path, "file:")
|
|
if path == "" || path == ":memory:" {
|
|
return nil
|
|
}
|
|
|
|
directory := filepath.Dir(path)
|
|
if directory == "." || directory == "" {
|
|
return nil
|
|
}
|
|
return os.MkdirAll(directory, 0o755)
|
|
}
|