db.go raw

   1  package db
   2  
   3  import (
   4  	"fmt"
   5  	"strings"
   6  
   7  	"gorm.io/driver/postgres"
   8  	"gorm.io/driver/sqlite"
   9  	"gorm.io/gorm"
  10  	gorm_logger "gorm.io/gorm/logger"
  11  
  12  	"github.com/getAlby/hub/db/migrations"
  13  	sqlite_wrapper "github.com/getAlby/hub/db/sqlite-wrapper"
  14  	"github.com/getAlby/hub/logger"
  15  )
  16  
  17  type Config struct {
  18  	URI        string
  19  	LogQueries bool
  20  	DriverName string
  21  }
  22  
  23  func NewDB(uri string, logDBQueries bool) (*gorm.DB, error) {
  24  	return NewDBWithConfig(&Config{
  25  		URI:        uri,
  26  		LogQueries: logDBQueries,
  27  		DriverName: "",
  28  	})
  29  }
  30  
  31  func NewDBWithConfig(cfg *Config) (*gorm.DB, error) {
  32  	gormConfig := &gorm.Config{
  33  		TranslateError: true,
  34  	}
  35  	if cfg.LogQueries {
  36  		gormConfig.Logger = gorm_logger.Default.LogMode(gorm_logger.Info)
  37  	}
  38  
  39  	var ret *gorm.DB
  40  
  41  	if IsPostgresURI(cfg.URI) {
  42  		pgConfig := postgres.Config{
  43  			DriverName: cfg.DriverName,
  44  			DSN:        cfg.URI,
  45  		}
  46  		var err error
  47  		ret, err = newPostgresDB(pgConfig, gormConfig)
  48  		if err != nil {
  49  			return nil, err
  50  		}
  51  	} else {
  52  		sqliteURI := cfg.URI
  53  
  54  		// apply pragma if we're not running the tests
  55  		if !strings.Contains(sqliteURI, "?mode=memory") {
  56  			// see https://github.com/mattn/go-sqlite3?tab=readme-ov-file#connection-string
  57  			// _txlock: avoid SQLITE_BUSY errors with _txlock=immediate
  58  			// _auto_vacuum: properly cleanup disk when deleting records with auto_vacuum=1
  59  			// _busy_timeout: avoid SQLITE_BUSY errors with 5 second lock timeout
  60  			// _journal_mode: enables write-ahead log so that your reads do not block writes and vice-versa.
  61  			// _synchronous: sqlite will sync less frequently and be more performant, still safe to use because of the enabled WAL mode
  62  			// _cache_size: 20MB memory cache
  63  			sqliteURI = sqliteURI + "?_txlock=immediate&_foreign_keys=1&_auto_vacuum=1&_busy_timeout=5000&_journal_mode=WAL&_synchronous=NORMAL&_cache_size=-20000"
  64  		}
  65  
  66  		driverName := sqlite_wrapper.Sqlite3WrapperDriverName
  67  		if cfg.DriverName != "" {
  68  			driverName = cfg.DriverName
  69  		}
  70  
  71  		sqliteConfig := sqlite.Config{
  72  			DriverName: driverName,
  73  			DSN:        sqliteURI,
  74  		}
  75  
  76  		var err error
  77  		ret, err = newSqliteDB(sqliteConfig, gormConfig)
  78  		if err != nil {
  79  			return nil, err
  80  		}
  81  	}
  82  
  83  	logger.Logger.WithField("db_backend", ret.Dialector.Name()).Debug("loaded database")
  84  
  85  	err := migrations.Migrate(ret)
  86  	if err != nil {
  87  		logger.Logger.WithError(err).Error("Failed to migrate")
  88  		return nil, err
  89  	}
  90  
  91  	return ret, nil
  92  }
  93  
  94  func newSqliteDB(sqliteConfig sqlite.Config, gormConfig *gorm.Config) (*gorm.DB, error) {
  95  	gormDB, err := gorm.Open(sqlite.New(sqliteConfig), gormConfig)
  96  	if err != nil {
  97  		return nil, err
  98  	}
  99  
 100  	return gormDB, nil
 101  }
 102  
 103  func newPostgresDB(pgConfig postgres.Config, gormConfig *gorm.Config) (*gorm.DB, error) {
 104  	gormDB, err := gorm.Open(postgres.New(pgConfig), gormConfig)
 105  	if err != nil {
 106  		return nil, err
 107  	}
 108  
 109  	return gormDB, nil
 110  }
 111  
 112  func Stop(db *gorm.DB) error {
 113  	sqlDB, err := db.DB()
 114  	if err != nil {
 115  		return fmt.Errorf("failed to get database connection: %w", err)
 116  	}
 117  
 118  	dbBackend := db.Dialector.Name()
 119  	logger.Logger.WithField("db_backend", dbBackend).Debug("shutting down database")
 120  	if dbBackend == "sqlite" {
 121  		err = db.Exec("PRAGMA wal_checkpoint(FULL)", nil).Error
 122  		if err != nil {
 123  			logger.Logger.WithError(err).Error("Failed to execute wal endpoint")
 124  		}
 125  	}
 126  
 127  	err = sqlDB.Close()
 128  	if err != nil {
 129  		return fmt.Errorf("failed to close database connection: %w", err)
 130  	}
 131  	return nil
 132  }
 133  
 134  func IsPostgresURI(uri string) bool {
 135  	return strings.HasPrefix(uri, "postgresql://") ||
 136  		strings.HasPrefix(uri, "postgres://") // Schema used by the "testdb" package.
 137  }
 138