package acceptance import ( "context" "database/sql" "fmt" "os" "strings" "testing" "time" _ "github.com/go-sql-driver/mysql" "github.com/redis/go-redis/v9" ) type FixtureLoader struct { cfg Config } func newFixtureLoader(c Config) *FixtureLoader { return &FixtureLoader{cfg: c} } func (l *FixtureLoader) SeedIfConfigured(t *testing.T, ctx context.Context) { t.Helper() if !l.hasDBConfig() { t.Log("staging DB config not provided; skipping seed loader") return } sqlBytes, err := os.ReadFile(l.cfg.SeedSQLPath) if err != nil { t.Fatalf("read seed SQL %q: %v", l.cfg.SeedSQLPath, err) } db, err := sql.Open("mysql", l.dsn(true)) if err != nil { t.Fatalf("open staging DB: %v", err) } defer func() { _ = db.Close() }() if err := db.PingContext(ctx); err != nil { t.Fatalf("ping staging DB: %v", err) } statements := splitSQL(string(sqlBytes)) for _, statement := range statements { statement = strings.ReplaceAll(statement, "{{RUN_ID}}", l.cfg.RunID) if _, err := db.ExecContext(ctx, statement); err != nil { t.Fatalf("execute seed SQL: %v; statement=%s", err, statement) } } } func (l *FixtureLoader) NodeSecret(ctx context.Context) (string, error) { if l.cfg.NodeSecret != "" { return l.cfg.NodeSecret, nil } if !l.hasDBConfig() { return "", nil } db, err := sql.Open("mysql", l.dsn(false)) if err != nil { return "", fmt.Errorf("open staging DB: %w", err) } defer func() { _ = db.Close() }() ctx, cancel := context.WithTimeout(ctx, 5*time.Second) defer cancel() var secret string err = db.QueryRowContext(ctx, "SELECT value FROM system WHERE category = 'server' AND `key` = 'NodeSecret' LIMIT 1").Scan(&secret) if err != nil { return "", fmt.Errorf("read NodeSecret from staging DB: %w", err) } return secret, nil } func (l *FixtureLoader) RedisClient() *redis.Client { if l.cfg.Redis.Addr == "" { return nil } return redis.NewClient(&redis.Options{ Addr: l.cfg.Redis.Addr, Password: l.cfg.Redis.Password, DB: l.cfg.Redis.DB, }) } func (l *FixtureLoader) hasDBConfig() bool { return l.cfg.DB.Host != "" && l.cfg.DB.User != "" && l.cfg.DB.Name != "" } func (l *FixtureLoader) dsn(multiStatements bool) string { params := "parseTime=true&timeout=5s&readTimeout=10s&writeTimeout=10s" if multiStatements { params += "&multiStatements=true" } return fmt.Sprintf("%s:%s@tcp(%s)/%s?%s", l.cfg.DB.User, l.cfg.DB.Password, l.cfg.DB.Host, l.cfg.DB.Name, params, ) } func splitSQL(script string) []string { parts := strings.Split(script, ";") statements := make([]string, 0, len(parts)) for _, part := range parts { lines := strings.Split(part, "\n") kept := make([]string, 0, len(lines)) for _, line := range lines { line = strings.TrimSpace(line) if line == "" || strings.HasPrefix(line, "--") { continue } kept = append(kept, line) } statement := strings.Join(kept, "\n") if statement == "" || strings.HasPrefix(statement, "--") { continue } statements = append(statements, statement) } return statements }