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 }