Files
chorus/migrations/mysql_integration_test.go
T

200 lines
8.4 KiB
Go

package migrations
import (
"context"
"database/sql"
"fmt"
"os"
"os/exec"
"strings"
"testing"
"time"
_ "github.com/go-sql-driver/mysql"
)
// This test resets only the explicitly designated disposable migration database.
// It is opt-in so normal package tests cannot erase a developer's local data.
func TestMVP1MigrationsUpDownUpMySQL(t *testing.T) {
if os.Getenv("CHORUS_RUN_MIGRATION_TESTS") != "1" {
t.Skip("set CHORUS_RUN_MIGRATION_TESTS=1 for the disposable MySQL migration database")
}
migrationURL := strings.TrimSpace(os.Getenv("CHORUS_MIGRATE_URL"))
dsn := strings.TrimSpace(os.Getenv("CHORUS_DSN"))
if migrationURL == "" || dsn == "" {
t.Fatal("CHORUS_MIGRATE_URL and CHORUS_DSN are required for migration integration tests")
}
if _, err := exec.LookPath("migrate"); err != nil {
t.Fatalf("migrate executable is required: %v", err)
}
db, err := sql.Open("mysql", dsn)
if err != nil {
t.Fatal(err)
}
defer db.Close()
ctx, cancel := context.WithTimeout(context.Background(), 90*time.Second)
defer cancel()
if err := db.PingContext(ctx); err != nil {
t.Fatalf("connect disposable migration database: %v", err)
}
var databaseName string
if err := db.QueryRowContext(ctx, `SELECT DATABASE()`).Scan(&databaseName); err != nil {
t.Fatalf("read disposable migration database name: %v", err)
}
if databaseName != "chorus_test" {
t.Fatalf("refusing destructive migration test against database %q; expected chorus_test", databaseName)
}
runMigrate := func(args ...string) {
t.Helper()
commandArgs := append([]string{"-path", ".", "-database", migrationURL}, args...)
command := exec.CommandContext(ctx, "migrate", commandArgs...)
output, err := command.CombinedOutput()
if err != nil && !strings.Contains(string(output), "no change") {
t.Fatalf("migrate %s: %v\n%s", strings.Join(args, " "), err, output)
}
}
resetDisposableSchema(t, ctx, db)
runMigrate("goto", "1")
seedMVP0MigrationFixture(t, ctx, db)
runMigrate("goto", "4")
assertTableExists(t, ctx, db, "provider_credentials", true)
assertTableExists(t, ctx, db, "route_pools", true)
assertTableExists(t, ctx, db, "sys_user", true)
var credentialCount int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM provider_credentials WHERE provider_id = 1 AND credential_version = 1 AND key_id = 'migration-test-key'`).Scan(&credentialCount); err != nil {
t.Fatal(err)
}
if credentialCount != 1 {
t.Fatalf("credential backfill count = %d, want 1", credentialCount)
}
var activeCredential sql.NullInt64
var oldEnvelope sql.NullString
if err := db.QueryRowContext(ctx, `SELECT active_credential_id, api_key_enc FROM providers WHERE id = 1`).Scan(&activeCredential, &oldEnvelope); err != nil {
t.Fatal(err)
}
if !activeCredential.Valid || oldEnvelope.Valid {
t.Fatalf("provider credential migration did not activate the new credential safely")
}
assertCount(t, ctx, db, `SELECT COUNT(*) FROM provider_model_capabilities WHERE provider_model_id = 1 AND capability = 'text'`, 1)
assertCount(t, ctx, db, `SELECT COUNT(*) FROM sys_role WHERE role_key = 'chorus_operator'`, 1)
assertCount(t, ctx, db, `SELECT COUNT(*) FROM sys_api WHERE handle = 'chorus.providers.credential.rotate'`, 1)
runMigrate("down", "3")
assertTableExists(t, ctx, db, "provider_credentials", false)
assertTableExists(t, ctx, db, "sys_user", false)
assertCount(t, ctx, db, `SELECT COUNT(*) FROM providers WHERE id = 1 AND api_key_enc IS NOT NULL`, 1)
assertCount(t, ctx, db, `SELECT COUNT(*) FROM generations WHERE id = 1`, 1)
runMigrate("goto", "4")
assertTableExists(t, ctx, db, "provider_credentials", true)
assertCount(t, ctx, db, `SELECT COUNT(*) FROM provider_credentials WHERE provider_id = 1 AND credential_version = 1`, 1)
assertCount(t, ctx, db, `SELECT COUNT(*) FROM sys_role WHERE role_key = 'chorus_operator'`, 1)
// Encrypted envelopes cannot be converted in SQL. Simulate the explicitly
// controlled pre-migration cleanup, then verify the plaintext migration in
// both directions without carrying a real credential through the test.
if _, err := db.ExecContext(ctx, `UPDATE providers SET auth_type = 'none', active_credential_id = NULL WHERE id = 1`); err != nil {
t.Fatal(err)
}
if _, err := db.ExecContext(ctx, `DELETE FROM provider_credentials WHERE provider_id = 1`); err != nil {
t.Fatal(err)
}
runMigrate("goto", "5")
assertColumnExists(t, ctx, db, "provider_credentials", "api_key", true)
assertColumnExists(t, ctx, db, "provider_credentials", "api_key_enc", false)
assertCount(t, ctx, db, `SELECT COUNT(*) FROM sys_api WHERE handle = 'chorus.providers.credential.get'`, 1)
runMigrate("down", "1")
assertColumnExists(t, ctx, db, "provider_credentials", "api_key", false)
assertColumnExists(t, ctx, db, "provider_credentials", "api_key_enc", true)
runMigrate("up")
assertColumnExists(t, ctx, db, "provider_credentials", "api_key", true)
}
func resetDisposableSchema(t *testing.T, ctx context.Context, db *sql.DB) {
t.Helper()
rows, err := db.QueryContext(ctx, `SELECT table_name FROM information_schema.tables WHERE table_schema = DATABASE()`)
if err != nil {
t.Fatalf("list disposable migration tables: %v", err)
}
var tables []string
for rows.Next() {
var table string
if err := rows.Scan(&table); err != nil {
rows.Close()
t.Fatalf("scan disposable migration table: %v", err)
}
tables = append(tables, table)
}
if err := rows.Close(); err != nil {
t.Fatalf("close disposable migration table rows: %v", err)
}
if _, err := db.ExecContext(ctx, `SET FOREIGN_KEY_CHECKS = 0`); err != nil {
t.Fatalf("disable disposable migration foreign keys: %v", err)
}
defer func() {
if _, err := db.ExecContext(context.Background(), `SET FOREIGN_KEY_CHECKS = 1`); err != nil {
t.Errorf("restore disposable migration foreign keys: %v", err)
}
}()
for _, table := range tables {
quoted := "`" + strings.ReplaceAll(table, "`", "``") + "`"
if _, err := db.ExecContext(ctx, `DROP TABLE `+quoted); err != nil {
t.Fatalf("drop disposable migration table %s: %v", table, err)
}
}
}
func seedMVP0MigrationFixture(t *testing.T, ctx context.Context, db *sql.DB) {
t.Helper()
statements := []string{
`INSERT INTO users (id, email, password_hash, display_name, status) VALUES (1, 'migration-user@chorus.invalid', 'synthetic', 'Migration User', 'active')`,
`INSERT INTO providers (id, slug, name, base_url, auth_type, api_key_enc, enabled) VALUES (1, 'migration-provider', 'Migration Provider', 'https://provider.invalid', 'bearer', JSON_OBJECT('version', 1, 'key_id', 'migration-test-key', 'nonce', 'AA==', 'ciphertext', 'AA=='), FALSE)`,
`INSERT INTO provider_models (id, provider_id, name, model_id, api_type, kind, extra_body, timeout_ms, weight, enabled) VALUES (1, 1, 'Migration Model', 'migration-model', 'chat', 'text', JSON_OBJECT(), 30000, 100, TRUE)`,
`INSERT INTO prompt_templates (id, template_key, kind, api_type, name, version, template_text, enabled) VALUES (1, 'migration-template', 'text', 'chat', 'Migration Template', 1, '{{.UserPrompt}}', TRUE)`,
`INSERT INTO generations (id, user_id, provider_model_id, kind, status, idempotency_key, user_prompt, rendered_prompt, attempts, attempt_count) VALUES (1, 1, 1, 'text', 'pending', 'migration-key', 'migration input', 'migration input', JSON_ARRAY(), 0)`,
}
for _, statement := range statements {
if _, err := db.ExecContext(ctx, statement); err != nil {
t.Fatalf("seed MVP-0 fixture: %v", err)
}
}
}
func assertTableExists(t *testing.T, ctx context.Context, db *sql.DB, table string, want bool) {
t.Helper()
var count int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = DATABASE() AND table_name = ?`, table).Scan(&count); err != nil {
t.Fatal(err)
}
if (count == 1) != want {
t.Fatalf("table %s exists = %t, want %t", table, count == 1, want)
}
}
func assertCount(t *testing.T, ctx context.Context, db *sql.DB, query string, want int) {
t.Helper()
var got int
if err := db.QueryRowContext(ctx, query).Scan(&got); err != nil {
t.Fatal(err)
}
if got != want {
t.Fatal(fmt.Sprintf("query count = %d, want %d", got, want))
}
}
func assertColumnExists(t *testing.T, ctx context.Context, db *sql.DB, table, column string, want bool) {
t.Helper()
var count int
if err := db.QueryRowContext(ctx, `SELECT COUNT(*) FROM information_schema.columns WHERE table_schema = DATABASE() AND table_name = ? AND column_name = ?`, table, column).Scan(&count); err != nil {
t.Fatal(err)
}
if (count == 1) != want {
t.Fatalf("column %s.%s exists = %t, want %t", table, column, count == 1, want)
}
}