test_db.go raw

   1  package db
   2  
   3  import (
   4  	"fmt"
   5  	"net/url"
   6  	"os"
   7  	"strings"
   8  	"testing"
   9  
  10  	"github.com/peterldowns/pgtestdb"
  11  	"gorm.io/gorm"
  12  
  13  	"github.com/getAlby/hub/db"
  14  	"github.com/getAlby/hub/logger"
  15  )
  16  
  17  const defaultTestDB = "test.db"
  18  
  19  func GetTestDatabaseURI() string {
  20  	ret := os.Getenv("TEST_DATABASE_URI")
  21  	if ret == "" {
  22  		// TODO: use in-memory DB, or a temporary file
  23  		return defaultTestDB
  24  	}
  25  
  26  	return ret
  27  }
  28  
  29  func NewDB(t *testing.T) (*gorm.DB, error) {
  30  	dbUri := GetTestDatabaseURI()
  31  
  32  	if dbUri == defaultTestDB {
  33  		//in case the file was not removed in the last run, remove it before starting the test
  34  		logger.Logger.WithField("uri", defaultTestDB).Info("removing test db")
  35  		os.Remove(defaultTestDB)
  36  	}
  37  
  38  	logger.Logger.WithField("uri", dbUri).Info("Creating new test DB with URI")
  39  	return NewDBWithURI(t, dbUri)
  40  }
  41  
  42  func NewDBWithURI(t *testing.T, uri string) (*gorm.DB, error) {
  43  	if db.IsPostgresURI(uri) {
  44  		parsedURI, err := url.Parse(uri)
  45  		if err != nil {
  46  			return nil, fmt.Errorf("failed to parse postgres DB URI: %w", err)
  47  		}
  48  
  49  		var user, password string
  50  		if userInfo := parsedURI.User; userInfo != nil {
  51  			user = userInfo.Username()
  52  			password, _ = userInfo.Password()
  53  		}
  54  
  55  		dbName := strings.TrimPrefix(parsedURI.Path, "/")
  56  
  57  		config := pgtestdb.Custom(t, pgtestdb.Config{
  58  			DriverName: "pgx",
  59  			Host:       parsedURI.Hostname(),
  60  			Port:       parsedURI.Port(),
  61  			User:       user,
  62  			Password:   password,
  63  			Database:   dbName,
  64  		}, pgtestdb.NoopMigrator{})
  65  
  66  		uri = config.URL()
  67  	}
  68  
  69  	return db.NewDBWithConfig(&db.Config{
  70  		URI:        uri,
  71  		LogQueries: true,
  72  	})
  73  }
  74  
  75  func CloseDB(d *gorm.DB) {
  76  	if err := db.Stop(d); err != nil {
  77  		panic("failed to close database: " + err.Error())
  78  	}
  79  
  80  	if GetTestDatabaseURI() == defaultTestDB {
  81  		logger.Logger.WithField("uri", defaultTestDB).Info("removing test db")
  82  		os.Remove(defaultTestDB)
  83  	}
  84  }
  85