90 lines
3.1 KiB
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:"))
|
|
}
|
|
}
|