dbtest.go 1.8 KB

12345678910111213141516171819202122232425262728293031323334353637383940414243444546474849505152535455565758596061626364656667
  1. // Package dbtest opens throwaway panel databases for tests. Migrating a new
  2. // SQLite file costs ~850ms under -race; copying a migrated template ~130ms.
  3. package dbtest
  4. import (
  5. "os"
  6. "path/filepath"
  7. "sync"
  8. "testing"
  9. "github.com/mhsanaei/3x-ui/v3/internal/config"
  10. "github.com/mhsanaei/3x-ui/v3/internal/database"
  11. )
  12. var migrated struct {
  13. once sync.Once
  14. data []byte
  15. err error
  16. }
  17. // InitDB opens a new, fully migrated panel database at path and closes it when
  18. // t ends. Reopen an existing file with database.InitDB instead.
  19. func InitDB(t testing.TB, path string) {
  20. t.Helper()
  21. if config.GetDBKind() != "postgres" {
  22. if _, err := os.Stat(path); err == nil {
  23. t.Fatalf("dbtest.InitDB would overwrite existing %s; reopen it with database.InitDB", path)
  24. }
  25. data, err := migratedTemplate()
  26. if err != nil {
  27. t.Fatalf("build template database: %v", err)
  28. }
  29. if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
  30. t.Fatalf("create database dir: %v", err)
  31. }
  32. if err := os.WriteFile(path, data, 0o600); err != nil {
  33. t.Fatalf("copy template database: %v", err)
  34. }
  35. }
  36. if err := database.InitDB(path); err != nil {
  37. t.Fatalf("InitDB: %v", err)
  38. }
  39. t.Cleanup(func() { _ = database.CloseDB() })
  40. }
  41. func migratedTemplate() ([]byte, error) {
  42. migrated.once.Do(func() {
  43. dir, err := os.MkdirTemp("", "xui-dbtest-")
  44. if err != nil {
  45. migrated.err = err
  46. return
  47. }
  48. defer os.RemoveAll(dir)
  49. path := filepath.Join(dir, "template.db")
  50. if err := database.InitDB(path); err != nil {
  51. migrated.err = err
  52. return
  53. }
  54. // Closing the last connection checkpoints the WAL into the main file.
  55. if err := database.CloseDB(); err != nil {
  56. migrated.err = err
  57. return
  58. }
  59. migrated.data, migrated.err = os.ReadFile(path)
  60. })
  61. return migrated.data, migrated.err
  62. }