package testusers import ( "context" "database/sql" "fmt" "net" "os" "path/filepath" "regexp" "strings" "testing" "time" "github.com/go-sql-driver/mysql" ) // This test never connects to the application database. It creates its own // uniquely named schema on an explicitly supplied loopback MySQL connection. func isolatedMySQL(t *testing.T) *sql.DB { t.Helper() dsn := os.Getenv("IM_TEST_MYSQL_DSN") if dsn == "" { t.Skip("set IM_TEST_MYSQL_DSN to enable isolated local MySQL integration tests") } cfg, err := mysql.ParseDSN(dsn) if err != nil { t.Fatal("invalid test MySQL DSN") } host, _, err := net.SplitHostPort(cfg.Addr) if err != nil || cfg.Net != "tcp" || (host != "127.0.0.1" && host != "localhost" && host != "::1") || cfg.DBName != "" { t.Fatal("integration tests require a loopback TCP DSN without a database name") } cfg.ParseTime = true admin, err := sql.Open("mysql", cfg.FormatDSN()) if err != nil { t.Fatal(err) } t.Cleanup(func() { admin.Close() }) database := fmt.Sprintf("im_fixture_test_%d", time.Now().UnixNano()) if !regexp.MustCompile(`^im_fixture_test_[0-9]+$`).MatchString(database) { t.Fatal("unsafe isolated database name") } if _, err = admin.Exec("CREATE DATABASE `" + database + "` CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci"); err != nil { t.Fatal(err) } // Only drop the exact schema successfully created by this test. t.Cleanup(func() { if _, err := admin.Exec("DROP DATABASE `" + database + "`"); err != nil { t.Errorf("cleanup isolated schema %s: %v", database, err) } }) cfg.DBName, cfg.MultiStatements = database, true db, err := sql.Open("mysql", cfg.FormatDSN()) if err != nil { t.Fatal(err) } t.Cleanup(func() { db.Close() }) for _, name := range []string{"001_users.sql", "002_social.sql", "028_test_users.sql"} { data, err := os.ReadFile(filepath.Join("..", "..", "migrations", name)) if err != nil { t.Fatal(err) } if _, err = db.Exec(string(data)); err != nil { t.Fatal(err) } } return db } func TestMySQLSeedRepeatabilityAndIsolation(t *testing.T) { db := isolatedMySQL(t) var database string if err := db.QueryRow(`SELECT DATABASE()`).Scan(&database); err != nil { t.Fatal(err) } if _, err := db.Exec(`INSERT INTO users (public_id,password_hash) VALUES ('REAL_FIXTURE','sentinel')`); err != nil { t.Fatal(err) } ctx := context.Background() if _, err := Seed(ctx, db, "https://example.com/uploads", "wrong_database"); err == nil { t.Fatal("missing target-database protection") } first, err := Seed(ctx, db, "https://example.com/uploads", database) if err != nil || first.Created != 100 || first.Male != 50 || first.Female != 50 || len(first.IDs) != 100 { t.Fatalf("first seed: %+v, %v", first, err) } second, err := Seed(ctx, db, "https://example.com/uploads", database) if err != nil || second.Created != 0 || second.Skipped != 100 || second.IDs[0] != first.IDs[0] { t.Fatalf("repeat seed: %+v, %v", second, err) } for query, want := range map[string]int{ `SELECT COUNT(*) FROM users WHERE is_test=0 AND test_batch='' AND password_hash='sentinel'`: 1, `SELECT COUNT(*) FROM users WHERE is_test=1 AND phone_hash IS NULL AND phone_cipher IS NULL AND password_hash='!TEST_PROFILE_NO_LOGIN'`: 100, `SELECT COUNT(*) FROM user_profiles WHERE last_active_at IS NULL AND is_vip=0 AND vip_level=0`: 100, `SELECT COUNT(*) FROM user_profiles WHERE gender=1`: 50, `SELECT COUNT(*) FROM user_profiles WHERE gender=2`: 50, `SELECT COUNT(*) FROM user_privacy_settings WHERE distance_visible=1`: 100, `SELECT COUNT(*) FROM user_location_states WHERE source='fixture' AND latitude IS NOT NULL AND longitude IS NOT NULL`: 100, `SELECT COUNT(*) FROM user_sessions`: 0, } { var got int if err := db.QueryRow(query).Scan(&got); err != nil || got != want { t.Errorf("query %s: got %d want %d err %v", query, got, want, err) } } if _, err := db.Exec(`UPDATE users SET test_batch='changed' WHERE public_id='TESTCN000100'`); err != nil { t.Fatal(err) } if _, err := Seed(ctx, db, "https://example.com/uploads", database); err == nil || !strings.Contains(err.Error(), "expected 0 or 100") { t.Fatalf("partial batch should not be silently repaired: %v", err) } } func TestMySQLSeedRollsBackOnRealUserCollision(t *testing.T) { db := isolatedMySQL(t) var database string if err := db.QueryRow(`SELECT DATABASE()`).Scan(&database); err != nil { t.Fatal(err) } // Force a late collision to prove earlier inserts in this batch roll back. if _, err := db.Exec(`INSERT INTO users (public_id,password_hash) VALUES ('TESTCN000099','real-user-sentinel')`); err != nil { t.Fatal(err) } if _, err := Seed(context.Background(), db, "https://example.com/uploads", database); err == nil { t.Fatal("real-user collision must fail") } var total, tests, profiles int _ = db.QueryRow(`SELECT COUNT(*),COALESCE(SUM(is_test),0) FROM users`).Scan(&total, &tests) _ = db.QueryRow(`SELECT COUNT(*) FROM user_profiles`).Scan(&profiles) if total != 1 || tests != 0 || profiles != 0 { t.Fatalf("partial seed survived rollback: users=%d tests=%d profiles=%d", total, tests, profiles) } }