Files
chorus/internal/seed/mysql_integration_test.go

90 lines
3.1 KiB
Go

package seed
import (
"context"
"database/sql"
"fmt"
"os"
"strings"
"testing"
"time"
_ "github.com/go-sql-driver/mysql"
)
func TestSeedUserUsernameIdempotencyMySQL(t *testing.T) {
if os.Getenv("CHORUS_RUN_SEED_TESTS") != "1" {
t.Skip("set CHORUS_RUN_SEED_TESTS=1 for the disposable MySQL seed database")
}
dsn := strings.TrimSpace(os.Getenv("CHORUS_DSN"))
if dsn == "" {
t.Fatal("CHORUS_DSN is required")
}
db, err := sql.Open("mysql", dsn)
if err != nil {
t.Fatal(err)
}
defer db.Close()
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
var databaseName string
if err := db.QueryRowContext(ctx, `SELECT DATABASE()`).Scan(&databaseName); err != nil {
t.Fatal(err)
}
expected := strings.TrimSpace(os.Getenv("CHORUS_MIGRATION_TEST_DATABASE"))
if expected == "" {
expected = "chorus_test"
}
if databaseName != expected {
t.Fatalf("refusing seed integration test against %q; expected %s", databaseName, expected)
}
defaults, err := LoadDefaults()
if err != nil {
t.Fatal(err)
}
var providerExisted int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM providers WHERE slug = ?`, defaults.Provider.Slug).Scan(&providerExisted); err != nil {
t.Fatal(err)
}
promptExisted := make(map[string]int, len(defaults.Prompts))
for _, prompt := range defaults.Prompts {
var count int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM prompt_templates WHERE template_key = ?`, prompt.Key).Scan(&count); err != nil {
t.Fatal(err)
}
promptExisted[prompt.Key] = count
}
suffix := time.Now().UnixNano()
username := fmt.Sprintf("seed-%d", suffix)
email := fmt.Sprintf("seed-%d@chorus.invalid", suffix)
options := Options{UserUsername: username, UserEmail: email, UserPassword: fmt.Sprintf("synthetic-%d", suffix), ProviderBaseURL: "https://provider.invalid/v1"}
t.Cleanup(func() {
_, _ = db.Exec(`DELETE FROM users WHERE username = ?`, username)
if providerExisted == 0 {
_, _ = db.Exec(`DELETE FROM provider_model_capabilities WHERE provider_model_id IN (SELECT id FROM provider_models WHERE provider_id IN (SELECT id FROM providers WHERE slug = ?))`, defaults.Provider.Slug)
_, _ = db.Exec(`DELETE FROM provider_models WHERE provider_id IN (SELECT id FROM providers WHERE slug = ?)`, defaults.Provider.Slug)
_, _ = db.Exec(`DELETE FROM providers WHERE slug = ?`, defaults.Provider.Slug)
}
for _, prompt := range defaults.Prompts {
if promptExisted[prompt.Key] == 0 {
_, _ = db.Exec(`DELETE FROM prompt_templates WHERE template_key = ?`, prompt.Key)
}
}
})
if err := Run(ctx, db, options); err != nil {
t.Fatal(err)
}
if err := Run(ctx, db, options); err != nil {
t.Fatal(err)
}
var count int
var passwordHash string
if err := db.QueryRowContext(ctx, `SELECT COUNT(*), MAX(password_hash) FROM users WHERE username = ? AND email = ?`, username, email).Scan(&count, &passwordHash); err != nil {
t.Fatal(err)
}
if count != 1 || !strings.HasPrefix(passwordHash, "bcrypt:v1:") {
t.Fatalf("seeded user count=%d hash versioned=%t", count, strings.HasPrefix(passwordHash, "bcrypt:v1:"))
}
}