diff --git a/cmd/claimer/main.go b/cmd/claimer/main.go index c338ccf7a..319d8fc68 100644 --- a/cmd/claimer/main.go +++ b/cmd/claimer/main.go @@ -127,30 +127,7 @@ func main() { "did", did, "handle", handle) continue } - b.WriteString("UPDATE domain_claims SET deleted = NULL WHERE domain = ") - b.WriteString(emit.Esc(handle)) - b.WriteString(" AND did = ") - b.WriteString(emit.Esc(did)) - b.WriteString(" AND deleted IS NOT NULL;\n") - - cutoff := "strftime('%Y-%m-%dT%H:%M:%SZ', 'now', '-30 days')" - b.WriteString("UPDATE domain_claims SET domain = ") - b.WriteString(emit.Esc(handle)) - b.WriteString(", deleted = NULL WHERE did = ") - b.WriteString(emit.Esc(did)) - b.WriteString(" AND domain != ") - b.WriteString(emit.Esc(handle)) - b.WriteString(" AND deleted IS NOT NULL AND deleted < ") - b.WriteString(cutoff) - b.WriteString(";\n") - b.WriteString("DELETE FROM domain_claims WHERE domain = ") - b.WriteString(emit.Esc(handle)) - b.WriteString(" AND did != ") - b.WriteString(emit.Esc(did)) - b.WriteString(" AND deleted IS NOT NULL AND deleted < ") - b.WriteString(cutoff) - b.WriteString(";\n") - + writeClaimSQL(&b, did, handle) rows = append(rows, []emit.Cell{{V: did}, {V: handle}, {Nil: true}}) logger.Info("would claim domain", "did", did, "domain", handle) } @@ -211,6 +188,33 @@ func main() { } } +func writeClaimSQL(b *strings.Builder, did, handle string) { + cutoff := "strftime('%Y-%m-%dT%H:%M:%SZ', 'now', '-30 days')" + b.WriteString("UPDATE domain_claims SET deleted = NULL WHERE domain = ") + b.WriteString(emit.Esc(handle)) + b.WriteString(" AND did = ") + b.WriteString(emit.Esc(did)) + b.WriteString(" AND deleted IS NOT NULL;\n") + + b.WriteString("DELETE FROM domain_claims WHERE domain = ") + b.WriteString(emit.Esc(handle)) + b.WriteString(" AND did != ") + b.WriteString(emit.Esc(did)) + b.WriteString(" AND deleted IS NOT NULL AND deleted < ") + b.WriteString(cutoff) + b.WriteString(";\n") + + b.WriteString("UPDATE domain_claims SET domain = ") + b.WriteString(emit.Esc(handle)) + b.WriteString(", deleted = NULL WHERE did = ") + b.WriteString(emit.Esc(did)) + b.WriteString(" AND domain != ") + b.WriteString(emit.Esc(handle)) + b.WriteString(" AND deleted IS NOT NULL AND deleted < ") + b.WriteString(cutoff) + b.WriteString(";\n") +} + func sitesDir() string { wd, err := os.Getwd() if err != nil { diff --git a/cmd/claimer/main_test.go b/cmd/claimer/main_test.go new file mode 100644 index 000000000..d5610037c --- /dev/null +++ b/cmd/claimer/main_test.go @@ -0,0 +1,53 @@ +package main + +import ( + "database/sql" + "strings" + "testing" + + _ "github.com/mattn/go-sqlite3" +) + +func TestClaimSQLFreesExpiredDomainBeforeRename(t *testing.T) { + var b strings.Builder + writeClaimSQL(&b, "did:plc:y", "bob.tngl.sh") + statements := strings.Split(strings.TrimSuffix(b.String(), ";\n"), ";\n") + if len(statements) != 3 { + t.Fatalf("claim operations = %d, want undelete, delete, rename", len(statements)) + } + setup := func() *sql.DB { + t.Helper() + db, err := sql.Open("sqlite3", ":memory:") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { db.Close() }) + _, err = db.Exec(`CREATE TABLE domain_claims ( + did TEXT NOT NULL UNIQUE, + domain TEXT NOT NULL UNIQUE, + deleted TEXT + ); + INSERT INTO domain_claims (did, domain, deleted) VALUES + ('did:plc:z', 'bob.tngl.sh', strftime('%Y-%m-%dT%H:%M:%SZ', 'now', '-40 days')), + ('did:plc:y', 'bob.tngl.io', strftime('%Y-%m-%dT%H:%M:%SZ', 'now', '-40 days'));`) + if err != nil { + t.Fatal(err) + } + return db + } + oldOrder := strings.Join([]string{statements[0], statements[2], statements[1]}, ";\n") + ";" + if _, err := setup().Exec(oldOrder); err == nil || !strings.Contains(err.Error(), "UNIQUE constraint failed") { + t.Fatalf("old rename-before-delete error = %v, want UNIQUE violation", err) + } + db := setup() + if _, err := db.Exec(b.String()); err != nil { + t.Fatalf("claim SQL after freeing the domain: %v", err) + } + var owner string + if err := db.QueryRow("SELECT did FROM domain_claims WHERE domain = 'bob.tngl.sh' AND deleted IS NULL").Scan(&owner); err != nil { + t.Fatal(err) + } + if owner != "did:plc:y" { + t.Fatalf("domain owner = %q, want did:plc:y", owner) + } +} diff --git a/cmd/sites-migrate/main.go b/cmd/sites-migrate/main.go index 781ae66d8..305542c95 100644 --- a/cmd/sites-migrate/main.go +++ b/cmd/sites-migrate/main.go @@ -21,11 +21,11 @@ import ( var ( dbPath = flag.String("db", "appview.db", "path to the appview SQLite database") outFile = flag.String("out", "/tmp/sites-migrate.sql", "where to write the emitted SQL") - validate = flag.Bool("validate", false, "compare source, emitted and D1 row counts, exiting non-zero on any mismatch") - apply = flag.Bool("apply", false, "execute the emitted SQL against D1 via wrangler d1 execute --file") + validate = flag.Bool("validate", false, "compare source, emitted and D1 row counts after -apply, or on an existing D1") + apply = flag.Bool("apply", false, "import into an empty D1 via wrangler d1 execute --file") remote = flag.Bool("remote", false, "with -apply, target the remote D1 database instead of the local one") local = flag.Bool("local", false, "target the local D1 database: -validate defaults to the remote one, and -local overrides -remote") - yes = flag.Bool("yes", false, "required with -apply: confirm that the emitted script wipes every row of all four tables") + yes = flag.Bool("yes", false, "required with -apply: confirm the target D1 import") d1Name = flag.String("d1", "tangled-sites", "D1 database name for wrangler d1 execute") wranglerDir = flag.String("wrangler-dir", "", "directory that owns the wrangler config; empty resolves to /sites") ) @@ -50,7 +50,7 @@ func main() { flag.Parse() if *apply && !*yes { - fmt.Fprintf(os.Stderr, "-apply requires -yes: the emitted script deletes every row of domain_claims, site_repos, repo_sites and site_deploys before loading the source snapshot\n") + fmt.Fprintln(os.Stderr, "-apply requires -yes: confirm the D1 import") os.Exit(1) } @@ -77,38 +77,53 @@ func main() { } logger.Info("wrote migration SQL", "path", *outFile, "bytes", len(sqlFile)) + targetRemote := *remote && !*local + if *apply { + if err := emptyImportTarget(func(table string) (int, error) { + return d1Count(ctx, *d1Name, table, *wranglerDir, !targetRemote) + }); err != nil { + logger.Error("refusing to overwrite an existing D1 database", "err", err) + os.Exit(1) + } + if err := applySQL(ctx, *outFile, *d1Name, *wranglerDir, targetRemote); err != nil { + logger.Error("applying migration SQL", "path", *outFile, "err", err) + os.Exit(1) + } + logger.Info("applied migration SQL", "path", *outFile, "d1", *d1Name, "remote", targetRemote) + } + if *validate { emitted, err := countEmitted(*outFile) if err != nil { logger.Error("counting emitted rows", "path", *outFile, "err", err) os.Exit(1) } - if err := checkParity(ctx, counts, emitted, *d1Name, *wranglerDir, *local); err != nil { + validateLocal := *local || (*apply && !targetRemote) + if err := checkParity(ctx, counts, emitted, *d1Name, *wranglerDir, validateLocal); err != nil { logger.Error("parity check failed", "err", err) os.Exit(1) } logger.Info("parity check passed") } - if *apply { - targetRemote := *remote && !*local - if err := applySQL(ctx, *outFile, *d1Name, *wranglerDir, targetRemote); err != nil { - logger.Error("applying migration SQL", "path", *outFile, "err", err) - os.Exit(1) +} + +func emptyImportTarget(count func(table string) (int, error)) error { + for _, in := range inserts { + n, err := count(in.table) + if err != nil { + return err + } + if n != 0 { + return fmt.Errorf("%s contains %d row(s); import requires an empty D1", in.table, n) } - logger.Info("applied migration SQL", "path", *outFile, "d1", *d1Name, "remote", targetRemote) } + return nil } func buildSQL(ctx context.Context, db *appviewdb.DB) (string, map[string]int, error) { var b strings.Builder - b.WriteString("-- WARNING: deletes every row of domain_claims, site_repos, repo_sites\n") - b.WriteString("-- and site_deploys before loading the appview snapshot below.\n") - b.WriteString("PRAGMA foreign_keys = OFF;\n") - b.WriteString("DELETE FROM site_deploys;\n") - b.WriteString("DELETE FROM repo_sites;\n") - b.WriteString("DELETE FROM site_repos;\n") - b.WriteString("DELETE FROM domain_claims;\n\n") + b.WriteString("-- Import into an empty D1; existing claims and site state must not be replaced.\n") counts := make(map[string]int, len(inserts)) for _, in := range inserts { diff --git a/cmd/sites-migrate/main_test.go b/cmd/sites-migrate/main_test.go new file mode 100644 index 000000000..abd228bc9 --- /dev/null +++ b/cmd/sites-migrate/main_test.go @@ -0,0 +1,43 @@ +package main + +import ( + "context" + "errors" + "path/filepath" + "strings" + "testing" + + appviewdb "tangled.org/core/appview/db" +) + +func TestImportScriptDoesNotEraseLiveClaims(t *testing.T) { + db, err := appviewdb.Make(context.Background(), filepath.Join(t.TempDir(), "appview.db")) + if err != nil { + t.Fatal(err) + } + defer db.Close() + + script, _, err := buildSQL(context.Background(), db) + if err != nil { + t.Fatal(err) + } + if strings.Contains(strings.ToUpper(script), "DELETE FROM") { + t.Fatal("migration script can erase the D1 backfill") + } +} + +func TestImportRequiresEmptyD1(t *testing.T) { + counts := map[string]int{"domain_claims": 1} + count := func(table string) (int, error) { return counts[table], nil } + if err := emptyImportTarget(count); err == nil || !strings.Contains(err.Error(), "domain_claims") { + t.Fatalf("existing backfilled claim accepted: %v", err) + } + counts["domain_claims"] = 0 + if err := emptyImportTarget(count); err != nil { + t.Fatalf("empty target rejected: %v", err) + } + boom := errors.New("D1 unavailable") + if err := emptyImportTarget(func(string) (int, error) { return 0, boom }); !errors.Is(err, boom) { + t.Fatalf("failed D1 preflight accepted: %v", err) + } +}