135 lines
5.5 KiB
Go
135 lines
5.5 KiB
Go
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)
|
|
}
|
|
}
|