Files
2026-08-04 15:34:34 +08:00

77 lines
1.9 KiB
Go

package testdb
import (
"fmt"
"net/url"
"os"
"strings"
"testing"
"time"
"affiliate_dash/internal/database"
"gorm.io/driver/postgres"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// New 使用真实 PostgreSQL 创建测试隔离 schema,避免 SQLite 与生产 SQL 行为不一致。
func New(t testing.TB, _ ...interface{}) *gorm.DB {
t.Helper()
dsn := os.Getenv("TEST_DATABASE_URL")
if dsn == "" {
dsn = os.Getenv("DATABASE_URL")
}
if dsn == "" {
dsn = "postgres://affiliate:affiliate_dev_password@127.0.0.1:15432/affiliate_dash?sslmode=disable"
}
cfg := &gorm.Config{Logger: logger.Default.LogMode(logger.Silent)}
adminDB, err := gorm.Open(postgres.Open(dsn), cfg)
if err != nil {
t.Fatalf("连接 PostgreSQL 测试库失败: %v", err)
}
adminSQL, err := adminDB.DB()
if err != nil {
t.Fatalf("获取 PostgreSQL 连接失败: %v", err)
}
schema := fmt.Sprintf("test_%d", time.Now().UnixNano())
if err := adminDB.Exec(`CREATE SCHEMA "` + schema + `"`).Error; err != nil {
t.Fatalf("创建测试 schema 失败: %v", err)
}
t.Cleanup(func() {
_ = adminDB.Exec(`DROP SCHEMA IF EXISTS "` + schema + `" CASCADE`).Error
_ = adminSQL.Close()
})
db, err := gorm.Open(postgres.Open(withSearchPath(dsn, schema)), cfg)
if err != nil {
t.Fatalf("连接测试 schema 失败: %v", err)
}
sqlDB, err := db.DB()
if err != nil {
t.Fatalf("获取测试 schema 连接失败: %v", err)
}
t.Cleanup(func() { _ = sqlDB.Close() })
if err := database.Migrate(db); err != nil {
t.Fatalf("迁移测试 schema 失败: %v", err)
}
return db
}
func withSearchPath(dsn, schema string) string {
parsed, err := url.Parse(dsn)
if err == nil && strings.HasPrefix(parsed.Scheme, "postgres") {
query := parsed.Query()
query.Set("search_path", schema)
parsed.RawQuery = query.Encode()
return parsed.String()
}
if strings.TrimSpace(dsn) == "" {
return dsn
}
return dsn + " search_path=" + schema
}