isolate.go 2.5 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889909192939495
  1. package testpg
  2. import (
  3. "context"
  4. "crypto/rand"
  5. "encoding/hex"
  6. "fmt"
  7. "net/url"
  8. "os"
  9. "strings"
  10. "time"
  11. "github.com/jackc/pgx/v5"
  12. "github.com/jackc/pgx/v5/pgxpool"
  13. )
  14. const (
  15. dbTypeEnv = "XUI_DB_TYPE"
  16. dbDSNEnv = "XUI_DB_DSN"
  17. )
  18. // IsolatePackage gives one test package its own PostgreSQL schema: package test
  19. // binaries run concurrently, and sharing public lets their migrations race.
  20. func IsolatePackage(packageName string) (func(), error) {
  21. if os.Getenv(dbTypeEnv) != "postgres" {
  22. return func() {}, nil
  23. }
  24. baseDSN := strings.TrimSpace(os.Getenv(dbDSNEnv))
  25. if baseDSN == "" {
  26. return func() {}, nil
  27. }
  28. suffix := make([]byte, 8)
  29. if _, err := rand.Read(suffix); err != nil {
  30. return nil, fmt.Errorf("generate PostgreSQL test schema suffix: %w", err)
  31. }
  32. schema := fmt.Sprintf("xui_%s_%d_%s", sanitize(packageName), os.Getpid(), hex.EncodeToString(suffix))
  33. ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
  34. defer cancel()
  35. admin, err := pgxpool.New(ctx, baseDSN)
  36. if err != nil {
  37. return nil, fmt.Errorf("open PostgreSQL test database: %w", err)
  38. }
  39. if _, err := admin.Exec(ctx, "CREATE SCHEMA "+pgx.Identifier{schema}.Sanitize()); err != nil {
  40. admin.Close()
  41. return nil, fmt.Errorf("create PostgreSQL test schema: %w", err)
  42. }
  43. isolatedDSN, err := withSearchPath(baseDSN, schema)
  44. if err != nil {
  45. admin.Close()
  46. return nil, err
  47. }
  48. if err := os.Setenv(dbDSNEnv, isolatedDSN); err != nil {
  49. admin.Close()
  50. return nil, fmt.Errorf("set isolated PostgreSQL test DSN: %w", err)
  51. }
  52. return func() {
  53. cleanupCtx, cleanupCancel := context.WithTimeout(context.Background(), 10*time.Second)
  54. defer cleanupCancel()
  55. _, _ = admin.Exec(cleanupCtx, "DROP SCHEMA "+pgx.Identifier{schema}.Sanitize()+" CASCADE")
  56. admin.Close()
  57. _ = os.Setenv(dbDSNEnv, baseDSN)
  58. }, nil
  59. }
  60. func withSearchPath(dsn, schema string) (string, error) {
  61. u, err := url.Parse(dsn)
  62. if err == nil && (u.Scheme == "postgres" || u.Scheme == "postgresql") {
  63. query := u.Query()
  64. query.Set("search_path", schema)
  65. u.RawQuery = query.Encode()
  66. return u.String(), nil
  67. }
  68. if strings.ContainsAny(schema, " '[]=\\") {
  69. return "", fmt.Errorf("unsafe PostgreSQL test schema name")
  70. }
  71. return strings.TrimSpace(dsn) + " search_path=" + schema, nil
  72. }
  73. func sanitize(value string) string {
  74. var result strings.Builder
  75. for _, r := range strings.ToLower(value) {
  76. if r >= 'a' && r <= 'z' || r >= '0' && r <= '9' || r == '_' {
  77. result.WriteRune(r)
  78. } else {
  79. result.WriteByte('_')
  80. }
  81. }
  82. if result.Len() == 0 {
  83. return "pkg"
  84. }
  85. return result.String()
  86. }