From bb6de4d0c1ec2ff40882602e13669132a4e9cec6 Mon Sep 17 00:00:00 2001 From: slZhong <1542123803@qq.com> Date: Tue, 11 Aug 2026 16:41:42 +0800 Subject: [PATCH 1/3] Add staging-safe logs slimming job --- Dockerfile.logs-slimming | 16 + cloudbuild.logs-slimming.yaml | 9 + cmd/logs_slimming/config.go | 193 +++++ cmd/logs_slimming/connection.go | 59 ++ cmd/logs_slimming/core_test.go | 615 +++++++++++++++ cmd/logs_slimming/cutover.go | 1267 +++++++++++++++++++++++++++++++ cmd/logs_slimming/ddl.go | 229 ++++++ cmd/logs_slimming/evidence.go | 89 +++ cmd/logs_slimming/main.go | 101 +++ cmd/logs_slimming/ops.go | 1083 ++++++++++++++++++++++++++ cmd/logs_slimming/sqlgen.go | 188 +++++ cmd/logs_slimming/topology.go | 91 +++ 12 files changed, 3940 insertions(+) create mode 100644 Dockerfile.logs-slimming create mode 100644 cloudbuild.logs-slimming.yaml create mode 100644 cmd/logs_slimming/config.go create mode 100644 cmd/logs_slimming/connection.go create mode 100644 cmd/logs_slimming/core_test.go create mode 100644 cmd/logs_slimming/cutover.go create mode 100644 cmd/logs_slimming/ddl.go create mode 100644 cmd/logs_slimming/evidence.go create mode 100644 cmd/logs_slimming/main.go create mode 100644 cmd/logs_slimming/ops.go create mode 100644 cmd/logs_slimming/sqlgen.go create mode 100644 cmd/logs_slimming/topology.go diff --git a/Dockerfile.logs-slimming b/Dockerfile.logs-slimming new file mode 100644 index 00000000000..18769cdd36d --- /dev/null +++ b/Dockerfile.logs-slimming @@ -0,0 +1,16 @@ +FROM golang:1.26.1-alpine@sha256:2389ebfa5b7f43eeafbd6be0c3700cc46690ef842ad962f6c5bd6be49ed82039 AS builder + +ENV CGO_ENABLED=0 GOOS=linux GOARCH=amd64 +WORKDIR /build + +COPY go.mod go.sum ./ +RUN go mod download + +COPY cmd/logs_slimming ./cmd/logs_slimming +RUN go build -trimpath -ldflags="-s -w" -o /out/logs-slimming ./cmd/logs_slimming + +FROM gcr.io/distroless/static-debian12:nonroot + +COPY --from=builder /out/logs-slimming /logs-slimming +USER nonroot:nonroot +ENTRYPOINT ["/logs-slimming"] diff --git a/cloudbuild.logs-slimming.yaml b/cloudbuild.logs-slimming.yaml new file mode 100644 index 00000000000..cbe031514b1 --- /dev/null +++ b/cloudbuild.logs-slimming.yaml @@ -0,0 +1,9 @@ +steps: + - name: gcr.io/cloud-builders/docker + args: + - build + - --file=Dockerfile.logs-slimming + - --tag=${_IMAGE} + - . +images: + - ${_IMAGE} diff --git a/cmd/logs_slimming/config.go b/cmd/logs_slimming/config.go new file mode 100644 index 00000000000..000bff2a29d --- /dev/null +++ b/cmd/logs_slimming/config.go @@ -0,0 +1,193 @@ +package main + +import ( + "errors" + "flag" + "fmt" + "os" + "regexp" + "slices" + "strconv" + "strings" + "time" +) + +const ( + dsnEnvironment = "LOGS_SLIMMING_DSN" + stagingSchema = "newapi_staging" + productionSchema = "newapi" +) + +var ( + identifierPattern = regexp.MustCompile(`^[a-z][a-z0-9_]{0,63}$`) + allowedCommands = []string{"preflight", "prepare", "backfill", "install-forward-trigger", "reconcile", "cutover", "rollback", "recover", "verify", "cleanup"} +) + +type config struct { + command string + dsn string + schema string + source string + target string + old string + checkpoint string + batch string + expectedProject string + expectedInstance string + expectedHostname string + expectedServerUUID string + expectedDatabaseUser string + channelIDs []int64 + batchSize int + batchDelay time.Duration + statementTimeout time.Duration + ddlTimeout time.Duration + lockWaitSeconds int + autoIncrementReserve uint64 + confirmCleanup string + evidencePath string + triggerDefiner string + phase string + upperBound int64 + maxThreadsRunning int +} + +func parseConfig(args []string) (config, error) { + if len(args) == 0 || !slices.Contains(allowedCommands, args[0]) { + return config{}, fmt.Errorf("first argument must be one of: %s", strings.Join(allowedCommands, ", ")) + } + cfg := config{command: args[0]} + fs := flag.NewFlagSet("logs_slimming "+cfg.command, flag.ContinueOnError) + fs.StringVar(&cfg.schema, "schema", stagingSchema, "database schema; staging is the default and only implicit target") + fs.StringVar(&cfg.source, "source", "logs", "active source table") + fs.StringVar(&cfg.target, "target", "", "owned shadow table") + fs.StringVar(&cfg.old, "old", "", "owned rollback table") + fs.StringVar(&cfg.checkpoint, "checkpoint-table", "", "owned durable checkpoint table") + fs.StringVar(&cfg.batch, "batch", "", "immutable migration batch identifier") + fs.StringVar(&cfg.expectedProject, "expected-project", "", "expected GCP project identity") + fs.StringVar(&cfg.expectedInstance, "expected-instance", "", "expected Cloud SQL instance identity") + fs.StringVar(&cfg.expectedHostname, "expected-hostname", "", "exact MySQL @@hostname") + fs.StringVar(&cfg.expectedServerUUID, "expected-server-uuid", "", "exact MySQL @@server_uuid") + fs.StringVar(&cfg.expectedDatabaseUser, "expected-db-user", "", "exact MySQL CURRENT_USER() in user@host form") + channels := fs.String("channel-ids", "", "comma-separated frozen Codex channel IDs") + fs.IntVar(&cfg.batchSize, "batch-size", 2000, "keyset rows per batch (1-5000)") + fs.DurationVar(&cfg.batchDelay, "batch-delay", 250*time.Millisecond, "delay between batches") + fs.DurationVar(&cfg.statementTimeout, "statement-timeout", 2*time.Second, "per data statement wall timeout") + fs.DurationVar(&cfg.ddlTimeout, "ddl-timeout", 3*time.Second, "DDL wall watchdog timeout") + fs.IntVar(&cfg.lockWaitSeconds, "lock-wait-seconds", 1, "MySQL metadata lock wait timeout (1-3)") + reserve := fs.Uint64("auto-increment-reserve", 1_000_000, "cutover AUTO_INCREMENT safety reserve") + fs.StringVar(&cfg.confirmCleanup, "confirm-cleanup", "", "must exactly equal ownership marker for cleanup") + fs.StringVar(&cfg.evidencePath, "evidence", "", "append-only JSONL evidence path; defaults to stdout") + fs.StringVar(&cfg.triggerDefiner, "trigger-definer", "", "durable MySQL account in user@host form") + fs.StringVar(&cfg.phase, "phase", "seed", "checkpoint phase: seed, gap, fresh, incremental, rollback-gap") + fs.Int64Var(&cfg.upperBound, "upper-bound", 0, "inclusive backfill/reconcile upper ID bound") + fs.IntVar(&cfg.maxThreadsRunning, "max-threads-running", 32, "fail closed when Threads_running exceeds this value") + if err := fs.Parse(args[1:]); err != nil { + return config{}, err + } + cfg.dsn = os.Getenv(dsnEnvironment) + cfg.autoIncrementReserve = *reserve + parsed, err := parseChannelIDs(*channels) + if err != nil { + return config{}, err + } + cfg.channelIDs = parsed + if cfg.target == "" && cfg.batch != "" { + cfg.target = "logs_compact_" + cfg.batch + } + if cfg.old == "" && cfg.batch != "" { + cfg.old = "logs_old_" + cfg.batch + } + if cfg.checkpoint == "" && cfg.batch != "" { + cfg.checkpoint = "logs_slim_checkpoint_" + cfg.batch + } + if err := cfg.validate(); err != nil { + return config{}, err + } + return cfg, nil +} + +func (c config) validate() error { + if !slices.Contains(allowedCommands, c.command) { + return fmt.Errorf("unsupported command %q", c.command) + } + if c.schema == productionSchema { + return errors.New("this staging artifact is hard-denied from the production schema; production requires a separately reviewed build") + } else if c.schema != stagingSchema { + return fmt.Errorf("schema %q is denied; this artifact only permits %q", c.schema, stagingSchema) + } + if !regexp.MustCompile(`^[a-z0-9_]{1,32}$`).MatchString(c.batch) { + return fmt.Errorf("unsafe or missing batch %q", c.batch) + } + for label, value := range map[string]string{ + "schema": c.schema, "source": c.source, "target": c.target, + "old": c.old, "checkpoint": c.checkpoint, + } { + if !identifierPattern.MatchString(value) { + return fmt.Errorf("unsafe or missing %s %q", label, value) + } + } + if c.source == c.target || c.source == c.old || c.target == c.old { + return errors.New("source, target, and old table names must be distinct") + } + if !strings.HasSuffix(c.target, c.batch) || !strings.HasSuffix(c.old, c.batch) || !strings.HasSuffix(c.checkpoint, c.batch) { + return errors.New("target, old, and checkpoint names must end with the immutable batch") + } + if c.expectedProject == "" || c.expectedInstance == "" || c.expectedHostname == "" || c.expectedServerUUID == "" || c.expectedDatabaseUser == "" { + return errors.New("expected project, instance, hostname, server UUID, and database user are all required") + } + if strings.ContainsAny(c.expectedDatabaseUser, "`'\";\\") || len(strings.Split(c.expectedDatabaseUser, "@")) != 2 { + return errors.New("expected-db-user must be an exact simple user@host value") + } + if len(c.channelIDs) == 0 { + return errors.New("at least one frozen channel ID is required") + } + if c.batchSize < 1 || c.batchSize > 5000 { + return errors.New("batch-size must be between 1 and 5000") + } + if c.batchDelay < 0 || c.statementTimeout <= 0 || c.ddlTimeout < time.Second { + return errors.New("timeouts must be positive and batch delay cannot be negative") + } + if c.lockWaitSeconds < 1 || c.lockWaitSeconds > 3 { + return errors.New("lock-wait-seconds must be between 1 and 3") + } + if c.autoIncrementReserve < 1_000_000 { + return errors.New("auto-increment-reserve cannot be less than 1000000") + } + if c.maxThreadsRunning < 1 || c.maxThreadsRunning > 256 { + return errors.New("max-threads-running must be between 1 and 256") + } + if !slices.Contains([]string{"seed", "gap", "fresh", "incremental", "ddl-intent", "rollback-gap", "rollback-ready", "rollback-intent", "rollback-reconcile"}, c.phase) { + return fmt.Errorf("unsupported phase %q", c.phase) + } + return nil +} + +func parseChannelIDs(raw string) ([]int64, error) { + seen := make(map[int64]struct{}) + var ids []int64 + for _, part := range strings.Split(raw, ",") { + part = strings.TrimSpace(part) + if part == "" { + continue + } + id, err := strconv.ParseInt(part, 10, 64) + if err != nil || id <= 0 { + return nil, fmt.Errorf("invalid channel ID %q", part) + } + if _, ok := seen[id]; ok { + continue + } + seen[id] = struct{}{} + ids = append(ids, id) + } + slices.Sort(ids) + return ids, nil +} + +func quoteIdentifier(value string) (string, error) { + if !identifierPattern.MatchString(value) { + return "", fmt.Errorf("unsafe SQL identifier %q", value) + } + return "`" + value + "`", nil +} diff --git a/cmd/logs_slimming/connection.go b/cmd/logs_slimming/connection.go new file mode 100644 index 00000000000..a105b5e33a9 --- /dev/null +++ b/cmd/logs_slimming/connection.go @@ -0,0 +1,59 @@ +package main + +import ( + "context" + "database/sql" + "fmt" + "strings" + "time" +) + +type sessionFacts struct { + database, hostname, serverUUID string + currentUser, authenticatedUser string + connectionID int64 + isolation, binlog, sqlMode string + autoIncrementMode int +} + +func openVerifiedConn(ctx context.Context, db *sql.DB, c config, role string) (*sql.Conn, sessionFacts, error) { + conn, err := db.Conn(ctx) + if err != nil { + return nil, sessionFacts{}, fmt.Errorf("open %s connection: %w", role, err) + } + facts, err := verifyPhysicalConn(ctx, conn, c, role) + if err != nil { + _ = conn.Close() + return nil, sessionFacts{}, err + } + return conn, facts, nil +} + +func verifyPhysicalConn(ctx context.Context, conn *sql.Conn, c config, role string) (sessionFacts, error) { + checkCtx, cancel := context.WithTimeout(ctx, minDuration(c.statementTimeout, time.Second)) + defer cancel() + var facts sessionFacts + err := conn.QueryRowContext(checkCtx, "SELECT DATABASE(),@@hostname,@@server_uuid,CURRENT_USER(),USER(),CONNECTION_ID(),@@transaction_isolation,@@binlog_format,@@innodb_autoinc_lock_mode,@@SESSION.sql_mode").Scan( + &facts.database, &facts.hostname, &facts.serverUUID, &facts.currentUser, &facts.authenticatedUser, + &facts.connectionID, &facts.isolation, &facts.binlog, &facts.autoIncrementMode, &facts.sqlMode, + ) + if err != nil { + return facts, fmt.Errorf("verify %s connection: %w", role, err) + } + authUser := strings.SplitN(facts.authenticatedUser, "@", 2)[0] + currentUser := strings.SplitN(facts.currentUser, "@", 2)[0] + if facts.database != c.schema || facts.hostname != c.expectedHostname || facts.serverUUID != c.expectedServerUUID || facts.currentUser != c.expectedDatabaseUser || authUser != currentUser { + return facts, fmt.Errorf("%s connection identity mismatch database=%q hostname=%q server_uuid=%q current_user=%q authenticated_user=%q", role, facts.database, facts.hostname, facts.serverUUID, facts.currentUser, facts.authenticatedUser) + } + if facts.isolation != "READ-COMMITTED" || facts.binlog != "ROW" || facts.autoIncrementMode != 2 { + return facts, fmt.Errorf("%s connection session mismatch isolation=%q binlog=%q autoinc_mode=%d", role, facts.isolation, facts.binlog, facts.autoIncrementMode) + } + return facts, nil +} + +func minDuration(a, b time.Duration) time.Duration { + if a < b { + return a + } + return b +} diff --git a/cmd/logs_slimming/core_test.go b/cmd/logs_slimming/core_test.go new file mode 100644 index 00000000000..e66d745bda8 --- /dev/null +++ b/cmd/logs_slimming/core_test.go @@ -0,0 +1,615 @@ +package main + +import ( + "context" + "database/sql" + "errors" + "slices" + "strings" + "testing" + "time" + + "github.com/go-sql-driver/mysql" +) + +func validTestConfig() config { + return config{ + command: "backfill", + schema: "newapi_staging", + source: "logs", + target: "logs_compact_20260811", + old: "logs_old_20260811", + checkpoint: "logs_slim_checkpoint_20260811", + batch: "20260811", + expectedProject: "vocai-gemini-prod", + expectedInstance: "newapi-mysql", + expectedHostname: "newapi-mysql-primary", + expectedServerUUID: "00000000-0000-0000-0000-000000000001", + expectedDatabaseUser: "newapi_staging_app@%", + triggerDefiner: "newapi_staging_app@%", + phase: "seed", + channelIDs: []int64{57, 61}, + batchSize: 2000, + batchDelay: 250 * time.Millisecond, + statementTimeout: 2 * time.Second, + ddlTimeout: 3 * time.Second, + lockWaitSeconds: 1, + autoIncrementReserve: 1_000_000, + maxThreadsRunning: 32, + } +} + +func TestConfigRequiresStagingIdentity(t *testing.T) { + cfg := validTestConfig() + if err := cfg.validate(); err != nil { + t.Fatal(err) + } + + for name, mutate := range map[string]func(*config){ + "schema": func(c *config) { c.schema = "newapi" }, + "project": func(c *config) { c.expectedProject = "" }, + "instance": func(c *config) { c.expectedInstance = "" }, + "hostname": func(c *config) { c.expectedHostname = "" }, + "uuid": func(c *config) { c.expectedServerUUID = "" }, + } { + t.Run(name, func(t *testing.T) { + bad := cfg + mutate(&bad) + if err := bad.validate(); err == nil { + t.Fatal("unsafe configuration unexpectedly accepted") + } + }) + } +} + +func TestStagingArtifactHardRejectsProduction(t *testing.T) { + cfg := validTestConfig() + cfg.schema = "newapi" + if err := cfg.validate(); err == nil { + t.Fatal("staging artifact accepted production") + } +} + +func TestProductionAuthorizationFlagsDoNotExist(t *testing.T) { + args := []string{"preflight", "--allow-production", "--schema=newapi_staging"} + if _, err := parseConfig(args); err == nil { + t.Fatal("deprecated production authorization flag accepted") + } +} + +func TestRollbackCommandIsExplicitAndStillStagingOnly(t *testing.T) { + cfg := validTestConfig() + cfg.command = "rollback" + if err := cfg.validate(); err != nil { + t.Fatalf("rollback command rejected: %v", err) + } + cfg.schema = productionSchema + if err := cfg.validate(); err == nil { + t.Fatal("rollback command accepted production schema") + } +} + +func TestAuditedStagingSchemaOnlySwapsLastTwoColumns(t *testing.T) { + staging := auditedStagingLogColumns() + if len(staging) != len(auditedLogColumns) { + t.Fatalf("staging schema has %d columns", len(staging)) + } + for i := 0; i < 19; i++ { + if staging[i] != auditedLogColumns[i] { + t.Fatalf("staging column %d changed", i+1) + } + } + if staging[19].name != "upstream_request_id" || staging[20].name != "other" { + t.Fatalf("unexpected staging tail: %+v", staging[19:]) + } +} + +func TestValidateChannelSnapshot(t *testing.T) { + if err := validateChannelSnapshot([]int64{57, 61}, []int64{57, 61}); err != nil { + t.Fatalf("matching snapshots rejected: %v", err) + } + if err := validateChannelSnapshot([]int64{57}, []int64{57, 61}); err == nil { + t.Fatal("mismatched snapshots accepted") + } +} + +func TestCutoverObservationTablesAreExplicit(t *testing.T) { + want := []string{"metadata_locks", "threads"} + if !slices.Equal(cutoverObservationTables, want) { + t.Fatalf("observation tables=%v want=%v", cutoverObservationTables, want) + } + for _, table := range want { + if !strings.Contains(cutoverObservationProbeSQL(table), "performance_schema."+table) { + t.Fatalf("missing observation table %s", table) + } + } +} + +func TestCutoverObservationAccessReportsAllDeniedTables(t *testing.T) { + var probed []string + err := probeCutoverObservationAccess(context.Background(), func(_ context.Context, table string) error { + probed = append(probed, table) + return &mysql.MySQLError{Number: 1142, Message: "SELECT command denied"} + }) + if !slices.Equal(probed, cutoverObservationTables) { + t.Fatalf("probed=%v want=%v", probed, cutoverObservationTables) + } + if err == nil { + t.Fatal("missing SELECT privileges unexpectedly accepted") + } + for _, table := range cutoverObservationTables { + if !strings.Contains(err.Error(), "performance_schema."+table) { + t.Fatalf("error does not include %s: %v", table, err) + } + } +} + +func TestCutoverObservationAccessPreservesProbeFailures(t *testing.T) { + err := probeCutoverObservationAccess(context.Background(), func(_ context.Context, table string) error { + if table == "metadata_locks" { + return &mysql.MySQLError{Number: 1142, Message: "SELECT command denied"} + } + return errors.New("connection reset") + }) + if err == nil { + t.Fatal("failed probes unexpectedly accepted") + } + for _, want := range []string{ + "SELECT denied on performance_schema.metadata_locks", + "probe execution failed: performance_schema.threads: connection reset", + } { + if !strings.Contains(err.Error(), want) { + t.Fatalf("error missing %q: %v", want, err) + } + } +} + +func TestAuditedColumnsAllowOnlyKnownLogCollations(t *testing.T) { + want := columnSpec{name: "content", columnType: "longtext", nullable: "YES", charset: "utf8mb4", collation: "utf8mb4_unicode_ci"} + staging := want + staging.collation = "utf8mb4_0900_ai_ci" + if !matchesAuditedColumn(staging, want) { + t.Fatal("known staging collation rejected") + } + unexpected := want + unexpected.collation = "utf8mb4_bin" + if matchesAuditedColumn(unexpected, want) { + t.Fatal("unexpected collation accepted") + } +} + +func TestIdentifiersAreStrict(t *testing.T) { + for _, value := range []string{"logs", "logs_compact_20260811"} { + if _, err := quoteIdentifier(value); err != nil { + t.Fatalf("safe identifier %q rejected: %v", value, err) + } + } + for _, value := range []string{"newapi.logs", "logs` DROP TABLE x", "UPPER", ""} { + if _, err := quoteIdentifier(value); err == nil { + t.Fatalf("unsafe identifier %q accepted", value) + } + } +} + +func TestCopySQLIsBoundedExplicitAndNullSafe(t *testing.T) { + cfg := validTestConfig() + query, err := buildCopySQL(cfg) + if err != nil { + t.Fatal(err) + } + for _, required := range []string{ + "FORCE INDEX (PRIMARY)", "id > ?", "id <= ?", "id <= ?", + "COALESCE(user_id, -1) = 1", "COALESCE(token_id, 0) > 0", + "COALESCE(channel_id, -1) IN (57,61)", "ON DUPLICATE KEY UPDATE id = `newapi_staging`.`logs_compact_20260811`.id", + } { + if !strings.Contains(query, required) { + t.Errorf("copy SQL missing %q\n%s", required, query) + } + } + if strings.Contains(query, "SELECT *") { + t.Fatal("copy SQL must not use SELECT *") + } + if strings.Contains(query, "UPDATE id = id") { + t.Fatal("copy SQL no-op update must qualify the target id") + } + for _, column := range logColumns { + if !strings.Contains(query, column) { + t.Errorf("copy SQL missing column %s", column) + } + } +} + +func TestRetainedAndFilteredPredicatesAreExactComplements(t *testing.T) { + filtered := filteredPredicate("l", []int64{57, 61}) + retained := retainedPredicate("l", []int64{57, 61}) + if retained != "NOT ("+filtered+")" { + t.Fatalf("predicates are not exact complements: filtered=%s retained=%s", filtered, retained) + } + for _, want := range []string{"l.user_id", "l.token_id", "l.channel_id", "IN (57,61)"} { + if !strings.Contains(filtered, want) { + t.Fatalf("filtered predicate missing %q: %s", want, filtered) + } + } +} + +func TestFullRowEqualityCoversAllColumnsWithBinaryText(t *testing.T) { + equal := rowEqualitySQL("s", "d") + if got := strings.Count(equal, "<=>"); got != len(logColumns) { + t.Fatalf("row equality compares %d fields, want %d: %s", got, len(logColumns), equal) + } + for _, column := range textColumns { + if !strings.Contains(equal, "BINARY s."+column+" <=> BINARY d."+column) { + t.Errorf("text column %s is not compared with binary semantics", column) + } + } +} + +func TestTriggerSQLUsesStrictInsertAndFrozenPredicate(t *testing.T) { + cfg := validTestConfig() + query, err := buildForwardTriggerSQL(cfg) + if err != nil { + t.Fatal(err) + } + if strings.Contains(strings.ToUpper(query), "IGNORE") || strings.Contains(strings.ToUpper(query), "ON DUPLICATE") { + t.Fatal("forward trigger must surface collisions") + } + if !strings.Contains(query, "IF NOT (") || !strings.Contains(query, "IN (57,61)") { + t.Fatalf("forward trigger does not use frozen retained predicate: %s", query) + } + for _, prefix := range []string{"NEW.id", "NEW.user_id", "NEW.upstream_request_id"} { + if !strings.Contains(query, prefix) { + t.Errorf("forward trigger missing %s", prefix) + } + } +} + +func TestCheckpointCASIsGenerationGuarded(t *testing.T) { + query, err := checkpointCASSQL(validTestConfig()) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(query, "generation = generation + 1") || !strings.Contains(query, "WHERE id = 1 AND generation = ?") { + t.Fatalf("checkpoint update is not CAS guarded: %s", query) + } +} + +func TestTopologyClassification(t *testing.T) { + tests := []struct { + name string + state objectTopology + status topologyStatus + }{ + {"pre", objectTopology{source: true, target: true}, topologyPreCutover}, + {"post", objectTopology{source: true, old: true}, topologyPostCutover}, + {"unknown-all", objectTopology{source: true, target: true, old: true}, topologyUnknown}, + {"unknown-none", objectTopology{}, topologyUnknown}, + {"unknown-missing-live", objectTopology{target: true}, topologyUnknown}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := classifyTopology(tt.state); got != tt.status { + t.Fatalf("got %s want %s", got, tt.status) + } + }) + } +} + +func TestVerifyTopologyFollowsCheckpointPhase(t *testing.T) { + if got := classifyForVerify(checkpoint{phase: "fresh"}); got != topologyPreCutover { + t.Fatalf("fresh checkpoint classified as %s", got) + } + if got := classifyForVerify(checkpoint{phase: "rollback-ready"}); got != topologyPostCutover { + t.Fatalf("rollback-ready checkpoint classified as %s", got) + } +} + +func TestPostCutoverOwnedTableIsLiveSource(t *testing.T) { + if got := ownedTableForTopology(validTestConfig(), topologyPostCutover); got != "logs" { + t.Fatalf("POST ownership checked on %q", got) + } + if got := ownedTableForTopology(validTestConfig(), topologyPreCutover); got != "logs_compact_20260811" { + t.Fatalf("PRE ownership checked on %q", got) + } +} + +func TestPostCutoverRecoveryNeverAllowsTriggerCycle(t *testing.T) { + for _, state := range []triggerTopology{ + {forward: true, reverse: true}, + {forward: true, reverse: true, updateGuard: true}, + } { + if _, err := recoveryPlan(topologyPostCutover, state); err == nil { + t.Fatal("recovery accepted simultaneous forward and reverse triggers") + } + } + + plan, err := recoveryPlan(topologyPostCutover, triggerTopology{forward: true}) + if err != nil { + t.Fatal(err) + } + want := []recoveryStep{stepRecordRollbackBase, stepDropForward, stepCreateReverse, stepReconcileRollbackGap} + if len(plan) != len(want) { + t.Fatalf("plan=%v want=%v", plan, want) + } + for i := range want { + if plan[i] != want[i] { + t.Fatalf("plan[%d]=%s want=%s", i, plan[i], want[i]) + } + } +} + +func TestCleanupRefusesUnknownOrPostCutoverTopology(t *testing.T) { + for _, status := range []topologyStatus{topologyUnknown, topologyPostCutover} { + if _, err := cleanupPlan(status, true); err == nil { + t.Fatalf("cleanup accepted topology %s", status) + } + } + plan, err := cleanupPlan(topologyPreCutover, true) + if err != nil { + t.Fatal(err) + } + if len(plan) == 0 || plan[len(plan)-1] != cleanupDropCheckpoint { + t.Fatalf("unexpected cleanup plan %v", plan) + } +} + +func TestDDLKillerRequiresExactConnectionAndStatement(t *testing.T) { + statement := "RENAME TABLE `logs` TO `logs_old`, `logs_compact` TO `logs`" + if !sameDDL(statement, " RENAME TABLE `logs` TO `logs_old`, `logs_compact` TO `logs` ") { + t.Fatal("whitespace-only difference should match") + } + if sameDDL(statement, "ALTER TABLE `logs` ADD COLUMN bad INT") { + t.Fatal("different DDL unexpectedly matched") + } +} + +func TestCutoverAndRollbackRenameStatementsAreSymmetric(t *testing.T) { + cfg := validTestConfig() + cutoverSQL := renameStatement(cfg) + rollbackSQL := rollbackStatement(cfg) + for _, want := range []string{"`newapi_staging`.`logs` TO `newapi_staging`.`logs_old_20260811`", "`newapi_staging`.`logs_compact_20260811` TO `newapi_staging`.`logs`"} { + if !strings.Contains(cutoverSQL, want) { + t.Fatalf("cutover SQL missing %q: %s", want, cutoverSQL) + } + } + for _, want := range []string{"`newapi_staging`.`logs` TO `newapi_staging`.`logs_compact_20260811`", "`newapi_staging`.`logs_old_20260811` TO `newapi_staging`.`logs`"} { + if !strings.Contains(rollbackSQL, want) { + t.Fatalf("rollback SQL missing %q: %s", want, rollbackSQL) + } + } + tagged := taggedDDL(cfg, rollbackOperation, "abc123", rollbackSQL) + if !strings.Contains(tagged, "operation="+rollbackOperation+" nonce=abc123") || !strings.HasSuffix(tagged, rollbackSQL) { + t.Fatalf("rollback DDL tag is not exact: %s", tagged) + } +} + +func TestEvidenceRedactsSecrets(t *testing.T) { + sink := &memoryEvidence{} + e := newEvidence(sink, []string{"user:secret@tcp(host)/db", "secret"}) + e.emit("failed", map[string]any{"error": errors.New("dial user:secret@tcp(host)/db failed"), "dsn": "secret"}) + output := sink.String() + if strings.Contains(output, "secret") || strings.Contains(output, "user:") { + t.Fatalf("evidence leaked secret: %s", output) + } + if !strings.Contains(output, "") { + t.Fatalf("evidence did not mark redaction: %s", output) + } +} + +func TestReserveRejectsOverflowAndUsesMinimum(t *testing.T) { + if got, err := targetAutoIncrement(100, 10); err != nil || got != 110 { + t.Fatalf("got=%d err=%v", got, err) + } + if _, err := targetAutoIncrement(^uint64(0)-5, 10); err == nil { + t.Fatal("overflow accepted") + } +} + +func TestCutoverUsesDoubleReserveBeforeFinalBarrier(t *testing.T) { + got, err := plannedTargetAutoIncrement(1_000, 1_000_000) + if err != nil || got != 2_001_000 { + t.Fatalf("got=%d err=%v", got, err) + } + if _, err := plannedTargetAutoIncrement(1, ^uint64(0)); err == nil { + t.Fatal("double reserve overflow accepted") + } +} + +func TestCutoverCheckpointRequiresFreshCompletedNoIntent(t *testing.T) { + ready := checkpoint{phase: "fresh", last: 200, final: sql.NullInt64{Int64: 200, Valid: true}} + if err := cutoverCheckpointReady(ready); err != nil { + t.Fatal(err) + } + for name, mutate := range map[string]func(*checkpoint){ + "phase": func(s *checkpoint) { s.phase = "gap" }, + "incomplete": func(s *checkpoint) { s.last = 199 }, + "no-final": func(s *checkpoint) { s.final.Valid = false }, + "operation": func(s *checkpoint) { s.ddlOperation = sql.NullString{String: cutoverOperation, Valid: true} }, + "nonce": func(s *checkpoint) { s.ddlNonce = sql.NullString{String: "nonce", Valid: true} }, + } { + t.Run(name, func(t *testing.T) { + state := ready + mutate(&state) + if err := cutoverCheckpointReady(state); err == nil { + t.Fatal("unsafe checkpoint accepted") + } + }) + } +} + +func TestMDLDrainAllowsOnlyBarrierAndExactRename(t *testing.T) { + if !mdlOwnerAllowed(10, "GRANTED", 10, 20) { + t.Fatal("barrier shared MDL rejected") + } + if !mdlOwnerAllowed(20, "GRANTED", 10, 20) || !mdlOwnerAllowed(20, "PENDING", 10, 20) { + t.Fatal("RENAME granted/pending MDL rejected") + } + for _, row := range []struct { + owner int64 + status string + }{{30, "GRANTED"}, {30, "PENDING"}, {10, "PENDING"}} { + if mdlOwnerAllowed(row.owner, row.status, 10, 20) { + t.Fatalf("foreign/invalid MDL accepted: %+v", row) + } + } +} + +func TestRenameStatementIsAtomicAndBarrierDoesNotUseLockTables(t *testing.T) { + statement := renameStatement(validTestConfig()) + for _, required := range []string{"RENAME TABLE", "`newapi_staging`.`logs` TO `newapi_staging`.`logs_old_20260811`", "`newapi_staging`.`logs_compact_20260811` TO `newapi_staging`.`logs`"} { + if !strings.Contains(statement, required) { + t.Fatalf("rename statement missing %q: %s", required, statement) + } + } + if strings.Contains(strings.ToUpper(statement), "LOCK TABLES") { + t.Fatal("cutover must not use LOCK TABLES") + } +} + +func TestPostTriggerPlanIsIdempotentAndNeverCycles(t *testing.T) { + plan, err := postTriggerPlan(triggerTopology{forward: true}) + if err != nil || len(plan) != 2 || plan[0] != postTriggerDropForward || plan[1] != postTriggerCreateReverse { + t.Fatalf("forward transition plan=%v err=%v", plan, err) + } + plan, err = postTriggerPlan(triggerTopology{}) + if err != nil || len(plan) != 1 || plan[0] != postTriggerCreateReverse { + t.Fatalf("missing reverse repair plan=%v err=%v", plan, err) + } + plan, err = postTriggerPlan(triggerTopology{reverse: true}) + if err != nil || len(plan) != 0 { + t.Fatalf("already stable reverse plan=%v err=%v", plan, err) + } + if _, err := postTriggerPlan(triggerTopology{forward: true, reverse: true}); err == nil { + t.Fatal("forward/reverse cycle accepted") + } +} + +func TestPostGuardPlanMovesOldGuardAndIsIdempotentOnLive(t *testing.T) { + plan, err := postGuardPlan(true, "logs_old_20260811", "logs", "logs_old_20260811") + if err != nil || len(plan) != 2 || plan[0] != postGuardDropOld || plan[1] != postGuardCreateLive { + t.Fatalf("old guard plan=%v err=%v", plan, err) + } + plan, err = postGuardPlan(false, "", "logs", "logs_old_20260811") + if err != nil || len(plan) != 1 || plan[0] != postGuardCreateLive { + t.Fatalf("missing guard plan=%v err=%v", plan, err) + } + plan, err = postGuardPlan(true, "logs", "logs", "logs_old_20260811") + if err != nil || len(plan) != 0 { + t.Fatalf("live guard plan=%v err=%v", plan, err) + } + if _, err := postGuardPlan(true, "foreign", "logs", "logs_old_20260811"); err == nil { + t.Fatal("foreign guard table accepted") + } +} + +func TestGuardSQLCanTargetPostCutoverLiveTable(t *testing.T) { + cfg := validTestConfig() + query, err := buildGuardTriggerSQL(cfg, "update", cfg.source) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(query, "BEFORE UPDATE ON `newapi_staging`.`logs`") { + t.Fatalf("guard is not attached to active logs: %s", query) + } + if guardsOnLive(triggerTopology{updateGuard: true, deleteGuard: true, updateGuardTable: cfg.old, deleteGuardTable: cfg.source}, cfg.source) { + t.Fatal("rollback-ready accepted a guard still attached to old") + } + if !guardsOnLive(triggerTopology{updateGuard: true, deleteGuard: true, updateGuardTable: cfg.source, deleteGuardTable: cfg.source}, cfg.source) { + t.Fatal("live guard topology rejected") + } +} + +func TestNextPostVerifyEndUsesSmallestValidBoundedEndpoint(t *testing.T) { + tests := []struct { + name string + old, source sql.NullInt64 + want sql.NullInt64 + }{ + {"both-old-first", sql.NullInt64{Int64: 10, Valid: true}, sql.NullInt64{Int64: 20, Valid: true}, sql.NullInt64{Int64: 10, Valid: true}}, + {"both-source-first", sql.NullInt64{Int64: 30, Valid: true}, sql.NullInt64{Int64: 20, Valid: true}, sql.NullInt64{Int64: 20, Valid: true}}, + {"only-old", sql.NullInt64{Int64: 10, Valid: true}, sql.NullInt64{}, sql.NullInt64{Int64: 10, Valid: true}}, + {"only-source", sql.NullInt64{}, sql.NullInt64{Int64: 20, Valid: true}, sql.NullInt64{Int64: 20, Valid: true}}, + {"neither", sql.NullInt64{}, sql.NullInt64{}, sql.NullInt64{}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := nextPostVerifyEnd(tt.old, tt.source); got != tt.want { + t.Fatalf("got=%v want=%v", got, tt.want) + } + }) + } +} + +func TestFullPostHighWaterIncludesEitherTable(t *testing.T) { + if got := maxInt64(100, 2_000_100); got != 2_000_100 { + t.Fatalf("source high-water omitted: %d", got) + } + if got := maxInt64(3_000_000, 2_000_100); got != 3_000_000 { + t.Fatalf("old high-water omitted: %d", got) + } +} + +func TestCheckpointPhaseTransitionsAreExplicitAndBoundaryGuarded(t *testing.T) { + seed := checkpoint{phase: "seed", last: 100, seed: 100} + last, final, needsForward, err := checkpointTransitionPlan(seed, "gap", 150) + if err != nil || last != 100 || !final.Valid || final.Int64 != 150 || !needsForward { + t.Fatalf("seed->gap plan got last=%d final=%v forward=%t err=%v", last, final, needsForward, err) + } + seed.last = 99 + if _, _, _, err := checkpointTransitionPlan(seed, "gap", 150); err == nil { + t.Fatal("incomplete seed was allowed to transition") + } + gap := checkpoint{phase: "gap", last: 150, final: sql.NullInt64{Int64: 150, Valid: true}} + last, _, _, err = checkpointTransitionPlan(gap, "fresh", 0) + if err != nil || last != 0 { + t.Fatalf("gap->fresh must reset to zero: last=%d err=%v", last, err) + } + if _, _, _, err := checkpointTransitionPlan(gap, "stable", 0); err == nil { + t.Fatal("illegal phase transition accepted") + } +} + +func TestVerifyRejectsTransitionalRollbackPhases(t *testing.T) { + for _, phase := range []string{"rollback-intent", "rollback-reconcile"} { + if got := classifyForVerify(checkpoint{phase: phase}); got != topologyUnknown { + t.Fatalf("transitional phase %s classified as %s", phase, got) + } + } +} + +func TestLockAndOwnershipBindMigrationDomain(t *testing.T) { + a := validTestConfig() + b := a + b.batch = "other" + b.target = "logs_compact_other" + b.old = "logs_old_other" + b.checkpoint = "logs_slim_checkpoint_other" + if advisoryLockName(a) != advisoryLockName(b) { + t.Fatal("different batches for the same source must contend on one advisory lock") + } + if ownershipMarker(a) == ownershipMarker(b) { + t.Fatal("ownership marker must bind concrete batch objects") + } + b = a + b.source = "other_logs" + if advisoryLockName(a) == advisoryLockName(b) { + t.Fatal("advisory lock must bind the concrete source") + } +} + +func TestAdvisoryLockNameFitsMySQLLimit(t *testing.T) { + cfg := validTestConfig() + name := advisoryLockName(cfg) + if len(name) > 64 { + t.Fatalf("lock name has %d characters: %s", len(name), name) + } + otherBatch := cfg + otherBatch.batch = "another_batch" + if advisoryLockName(otherBatch) != name { + t.Fatal("same source migration domain produced a batch-specific lock") + } + otherSource := cfg + otherSource.source = "logs_other" + if advisoryLockName(otherSource) == name { + t.Fatal("different source tables share an advisory lock") + } +} diff --git a/cmd/logs_slimming/cutover.go b/cmd/logs_slimming/cutover.go new file mode 100644 index 00000000000..9d7e715a176 --- /dev/null +++ b/cmd/logs_slimming/cutover.go @@ -0,0 +1,1267 @@ +package main + +import ( + "context" + "crypto/sha256" + "database/sql" + "fmt" + "math" + "regexp" + "strconv" + "strings" + "time" +) + +const cutoverOperation = "RENAME_TABLE" +const rollbackOperation = "ROLLBACK_RENAME_TABLE" + +var autoIncrementPattern = regexp.MustCompile(`(?i)\bAUTO_INCREMENT=(\d+)\b`) + +func cutover(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + state, err := assertCutoverPreconditions(ctx, db, c) + if err != nil { + return err + } + sourceNext, err := showCreateAutoIncrement(ctx, db, c.schema, c.source) + if err != nil { + return err + } + targetNext, err := plannedTargetAutoIncrement(sourceNext, c.autoIncrementReserve) + if err != nil { + return err + } + if err := setTargetAutoIncrement(ctx, db, c, targetNext); err != nil { + return err + } + nonce, err := ddlNonce() + if err != nil { + return err + } + state, err = persistDDLIntent(ctx, db, c, state, nonce) + if err != nil { + return err + } + status, resolved, renameErr := renameWithMDLBarrier(ctx, db, c, ev, nonce, cutoverOperation, renameStatement(c), c.target) + switch status { + case topologyPreCutover: + if !resolved { + return fmt.Errorf("rename remained PRE but execution is unresolved; persisted DDL intent retained: %w", renameErr) + } + clearErr := clearPreCutoverIntent(ctx, db, c, state) + if renameErr != nil { + return fmt.Errorf("rename remained PRE: %w (intent cleanup: %v)", renameErr, clearErr) + } + if clearErr != nil { + return clearErr + } + return fmt.Errorf("rename remained PRE without a driver error") + case topologyPostCutover: + if err := ev.emit("cutover_topology_post", map[string]any{"nonce": nonce, "rename_error": renameErr}); err != nil { + return err + } + return stabilizePostCutover(ctx, db, c, ev) + default: + return fmt.Errorf("rename outcome UNKNOWN; manual topology diagnosis required: %w", renameErr) + } +} + +// rollback atomically restores the original full logs table as the live table. +// The compact table remains owned and receives retained rows through the +// forward trigger, so a later cutover can be retried without rebuilding it. +func rollback(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if err := assertRollbackReady(ctx, db, c, state); err != nil { + return err + } + if err := assertRuntimeSafe(ctx, db, c, state, false); err != nil { + return err + } + if err := fullPostVerify(ctx, db, c); err != nil { + return err + } + sourceNext, err := showCreateAutoIncrement(ctx, db, c.schema, c.source) + if err != nil { + return err + } + oldNext, err := plannedTargetAutoIncrement(sourceNext, c.autoIncrementReserve) + if err != nil { + return err + } + if err := setTableAutoIncrement(ctx, db, c, c.old, oldNext); err != nil { + return err + } + source, _ := qualified(c.schema, c.source) + var base sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+source).Scan(&base); err != nil { + return err + } + nonce, err := ddlNonce() + if err != nil { + return err + } + state, err = persistRollbackIntent(ctx, db, c, state, nonce, base.Int64) + if err != nil { + return err + } + status, resolved, renameErr := renameWithMDLBarrier(ctx, db, c, ev, nonce, rollbackOperation, rollbackStatement(c), c.old) + switch status { + case topologyPostCutover: + if !resolved { + return fmt.Errorf("rollback RENAME remained POST but execution is unresolved; persisted DDL intent retained: %w", renameErr) + } + clearErr := clearPostCutoverRollbackIntent(ctx, db, c, state) + if renameErr != nil { + return fmt.Errorf("rollback RENAME remained POST: %w (intent cleanup: %v)", renameErr, clearErr) + } + if clearErr != nil { + return clearErr + } + return fmt.Errorf("rollback RENAME remained POST without a driver error") + case topologyPreCutover: + if err := ev.emit("rollback_topology_pre", map[string]any{"nonce": nonce, "rename_error": renameErr}); err != nil { + return err + } + return stabilizePreRollback(ctx, db, c, ev) + default: + return fmt.Errorf("rollback RENAME outcome UNKNOWN; persisted intent retained: %w", renameErr) + } +} + +func rollbackStatement(c config) string { + source, _ := qualified(c.schema, c.source) + target, _ := qualified(c.schema, c.target) + old, _ := qualified(c.schema, c.old) + return "RENAME TABLE " + source + " TO " + target + ", " + old + " TO " + source +} + +func persistRollbackIntent(ctx context.Context, db *sql.DB, c config, state checkpoint, nonce string, base int64) (checkpoint, error) { + cp, _ := qualified(c.schema, c.checkpoint) + result, err := db.ExecContext(ctx, "UPDATE "+cp+" SET phase='rollback-intent',last_completed_end_id=?,rollback_base_id=?,ddl_operation=?,ddl_nonce=?,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='rollback-ready' AND generation=? AND ddl_operation IS NULL AND ddl_nonce IS NULL", base, base, rollbackOperation, nonce, state.generation) + if err != nil { + return state, err + } + affected, _ := result.RowsAffected() + if affected != 1 { + return state, fmt.Errorf("rollback intent CAS conflict generation=%d", state.generation) + } + state.phase = "rollback-intent" + state.last = base + state.rollback = sql.NullInt64{Int64: base, Valid: true} + state.generation++ + state.ddlOperation = sql.NullString{String: rollbackOperation, Valid: true} + state.ddlNonce = sql.NullString{String: nonce, Valid: true} + return state, nil +} + +func clearPostCutoverRollbackIntent(ctx context.Context, db *sql.DB, c config, state checkpoint) error { + cp, _ := qualified(c.schema, c.checkpoint) + result, err := db.ExecContext(ctx, "UPDATE "+cp+" SET phase='rollback-ready',ddl_operation=NULL,ddl_nonce=NULL,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='rollback-intent' AND generation=? AND ddl_operation=? AND ddl_nonce=?", state.generation, rollbackOperation, state.ddlNonce.String) + if err != nil { + return err + } + affected, _ := result.RowsAffected() + if affected != 1 { + return fmt.Errorf("clear rollback POST intent CAS conflict generation=%d", state.generation) + } + return nil +} + +func plannedTargetAutoIncrement(sourceNext, reserve uint64) (uint64, error) { + if reserve > math.MaxUint64/2 { + return 0, fmt.Errorf("AUTO_INCREMENT double reserve overflow") + } + return targetAutoIncrement(sourceNext, 2*reserve) +} + +func assertCutoverPreconditions(ctx context.Context, db *sql.DB, c config) (checkpoint, error) { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return state, err + } + if err := cutoverCheckpointReady(state); err != nil { + return state, fmt.Errorf("cutover requires fresh completed checkpoint with no DDL intent: phase=%s last=%d final=%v", state.phase, state.last, state.final) + } + if err := assertOwned(ctx, db, c, c.target); err != nil { + return state, err + } + if err := assertFrozenTableFingerprints(ctx, db, c, state); err != nil { + return state, err + } + o, triggers, err := observeTopology(ctx, db, c) + if err != nil { + return state, err + } + if classifyTopology(o) != topologyPreCutover || !triggers.forward || !triggers.updateGuard || !triggers.deleteGuard || triggers.reverse { + return state, fmt.Errorf("cutover trigger/topology precondition failed objects=%+v triggers=%+v", o, triggers) + } + if err := assertRuntimeSafe(ctx, db, c, state, true); err != nil { + return state, err + } + return state, nil +} + +func cutoverCheckpointReady(state checkpoint) error { + if state.phase != "fresh" || !state.final.Valid || state.last < state.final.Int64 || state.ddlOperation.Valid || state.ddlNonce.Valid { + return fmt.Errorf("checkpoint is not cutover-ready") + } + return nil +} + +func persistDDLIntent(ctx context.Context, db *sql.DB, c config, state checkpoint, nonce string) (checkpoint, error) { + cp, _ := qualified(c.schema, c.checkpoint) + result, err := db.ExecContext(ctx, "UPDATE "+cp+" SET phase='ddl-intent',ddl_operation=?,ddl_nonce=?,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='fresh' AND generation=? AND ddl_operation IS NULL AND ddl_nonce IS NULL", cutoverOperation, nonce, state.generation) + if err != nil { + return state, err + } + affected, _ := result.RowsAffected() + if affected != 1 { + return state, fmt.Errorf("DDL intent CAS conflict generation=%d", state.generation) + } + state.phase, state.generation = "ddl-intent", state.generation+1 + state.ddlOperation = sql.NullString{String: cutoverOperation, Valid: true} + state.ddlNonce = sql.NullString{String: nonce, Valid: true} + return state, nil +} + +func clearPreCutoverIntent(ctx context.Context, db *sql.DB, c config, state checkpoint) error { + cp, _ := qualified(c.schema, c.checkpoint) + result, err := db.ExecContext(ctx, "UPDATE "+cp+" SET phase='fresh',ddl_operation=NULL,ddl_nonce=NULL,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='ddl-intent' AND generation=? AND ddl_operation=? AND ddl_nonce=?", state.generation, state.ddlOperation.String, state.ddlNonce.String) + if err != nil { + return err + } + affected, _ := result.RowsAffected() + if affected != 1 { + return fmt.Errorf("clear PRE intent CAS conflict generation=%d", state.generation) + } + return nil +} + +func setTargetAutoIncrement(ctx context.Context, db *sql.DB, c config, next uint64) error { + return setTableAutoIncrement(ctx, db, c, c.target, next) +} + +func setTableAutoIncrement(ctx context.Context, db *sql.DB, c config, table string, next uint64) error { + qualifiedTable, _ := qualified(c.schema, table) + statement := fmt.Sprintf("ALTER TABLE %s AUTO_INCREMENT=%d", qualifiedTable, next) + observer := func(ctx context.Context, conn *sql.Conn, _ string) (ddlState, error) { + got, err := showCreateAutoIncrement(ctx, conn, c.schema, table) + if err != nil { + return ddlUnknown, err + } + if got >= next { + return ddlPost, nil + } + return ddlPre, nil + } + if err := runDDL(ctx, db, c, statement, observer); err != nil { + return err + } + got, err := showCreateAutoIncrement(ctx, db, c.schema, table) + if err != nil || got < next { + return fmt.Errorf("SHOW CREATE AUTO_INCREMENT readback failed got=%d want>=%d err=%v", got, next, err) + } + return nil +} + +type showCreateQueryer interface { + QueryRowContext(context.Context, string, ...any) *sql.Row +} + +func showCreateAutoIncrement(ctx context.Context, q showCreateQueryer, schema, table string) (uint64, error) { + qualifiedTable, err := qualified(schema, table) + if err != nil { + return 0, err + } + var name, create string + if err := q.QueryRowContext(ctx, "SHOW CREATE TABLE "+qualifiedTable).Scan(&name, &create); err != nil { + return 0, err + } + match := autoIncrementPattern.FindStringSubmatch(create) + if len(match) != 2 { + return 0, fmt.Errorf("SHOW CREATE TABLE %s lacks AUTO_INCREMENT", table) + } + return strconv.ParseUint(match[1], 10, 64) +} + +func renameStatement(c config) string { + source, _ := qualified(c.schema, c.source) + target, _ := qualified(c.schema, c.target) + old, _ := qualified(c.schema, c.old) + return "RENAME TABLE " + source + " TO " + old + ", " + target + " TO " + source +} + +// renameWithMDLBarrier uses ordinary reads in a short READ COMMITTED transaction. +// Their shared metadata locks block only the RENAME; LOCK TABLES is deliberately +// not used. Once the exact RENAME is the sole pending/granted foreign MDL holder, +// the final ID/AUTO_INCREMENT invariant is checked and COMMIT releases the barrier. +func renameWithMDLBarrier(ctx context.Context, db *sql.DB, c config, ev *evidence, nonce, operation, statement, futureActive string) (topologyStatus, bool, error) { + deadline := time.Now().Add(c.ddlTimeout) + opCtx, cancel := context.WithDeadline(ctx, deadline) + defer cancel() + barrier, barrierFacts, err := openVerifiedConn(opCtx, db, c, "cutover-barrier") + if err != nil { + return topologyUnknown, true, err + } + defer barrier.Close() + tx, err := barrier.BeginTx(opCtx, &sql.TxOptions{Isolation: sql.LevelReadCommitted, ReadOnly: true}) + if err != nil { + return topologyUnknown, true, err + } + barrierOpen := true + defer func() { + if barrierOpen { + _ = tx.Rollback() + } + }() + source, _ := qualified(c.schema, c.source) + future, _ := qualified(c.schema, futureActive) + for _, table := range []string{source, future} { + if _, err := tx.ExecContext(opCtx, "SELECT id FROM "+table+" LIMIT 0"); err != nil { + return topologyUnknown, true, fmt.Errorf("acquire shared MDL barrier on %s: %w", table, err) + } + } + ddlConn, ddlFacts, err := openVerifiedConn(opCtx, db, c, "cutover-rename") + if err != nil { + return topologyUnknown, true, err + } + control, _, err := openVerifiedConn(opCtx, db, c, "cutover-control") + if err != nil { + _ = ddlConn.Close() + return topologyUnknown, true, err + } + defer control.Close() + if _, err := ddlConn.ExecContext(opCtx, fmt.Sprintf("SET SESSION lock_wait_timeout=%d", c.lockWaitSeconds)); err != nil { + _ = ddlConn.Close() + return topologyUnknown, true, err + } + tagged := fmt.Sprintf("/*logs_slim batch=%s operation=%s nonce=%s*/ %s", c.batch, operation, nonce, statement) + if err := ev.emit("rename_barrier_intent", map[string]any{"nonce": nonce, "connection_id": ddlFacts.connectionID, "statement_sha256": fmt.Sprintf("%x", sha256.Sum256([]byte(statement)))}); err != nil { + return topologyUnknown, true, err + } + done := make(chan error, 1) + go func() { + _, execErr := ddlConn.ExecContext(opCtx, tagged) + done <- execErr + }() + abort := func(cause error) (topologyStatus, bool, error) { + killCtx, stop := boundedBackground(deadline, 300*time.Millisecond) + killErr := killExactDDL(killCtx, control, c, ctx, ddlFacts, tagged) + stop() + cancel() + rollbackErr := tx.Rollback() + barrierOpen = false + resolved := false + select { + case <-done: + resolved = true + case <-time.After(ddlSettleGrace): + } + if resolved { + _ = ddlConn.Close() + } else { + go reapDDLConn(done, ddlConn) + } + observeDeadline := maxTime(deadline, time.Now().Add(ddlObservationGrace)) + status, topologyErr := observeStableObjectTopology(db, c, observeDeadline) + return status, resolved, fmt.Errorf("%w; exact_kill=%v barrier_rollback=%v topology=%v", cause, killErr, rollbackErr, topologyErr) + } + pendingDeadline := deadline.Add(-600 * time.Millisecond) + if pendingDeadline.Before(time.Now()) { + pendingDeadline = time.Now().Add(50 * time.Millisecond) + } + pendingCtx, stopPending := context.WithDeadline(opCtx, pendingDeadline) + err = waitForExactPendingRename(pendingCtx, control, c, barrierFacts, ddlFacts, tagged) + stopPending() + if err != nil { + return abort(fmt.Errorf("pending RENAME proof failed: %w", err)) + } + var sourceMax, futureMax uint64 + if err := tx.QueryRowContext(opCtx, "SELECT COALESCE(MAX(id),0) FROM "+source).Scan(&sourceMax); err != nil { + return abort(fmt.Errorf("final source MAX: %w", err)) + } + if err := tx.QueryRowContext(opCtx, "SELECT COALESCE(MAX(id),0) FROM "+future).Scan(&futureMax); err != nil { + return abort(fmt.Errorf("final future-active MAX: %w", err)) + } + futureNext, err := showCreateAutoIncrement(opCtx, tx, c.schema, futureActive) + if err != nil { + return abort(fmt.Errorf("final future-active SHOW CREATE: %w", err)) + } + sourceFinalNext, err := safeNext(sourceMax) + if err != nil { + return abort(err) + } + required, err := targetAutoIncrement(sourceFinalNext, c.autoIncrementReserve) + if err != nil || futureNext < required || futureNext <= futureMax { + return abort(fmt.Errorf("final AUTO_INCREMENT invariant failed source_max=%d future_max=%d future_next=%d required=%d err=%v", sourceMax, futureMax, futureNext, required, err)) + } + if err := tx.Commit(); err != nil { + return abort(fmt.Errorf("release MDL barrier: %w", err)) + } + barrierOpen = false + var execErr error + resolved := false + select { + case execErr = <-done: + resolved = true + _ = ddlConn.Close() + case <-opCtx.Done(): + killCtx, stop := boundedBackground(deadline, 250*time.Millisecond) + killErr := killExactDDL(killCtx, control, c, ctx, ddlFacts, tagged) + stop() + execErr = fmt.Errorf("rename watchdog expired: %w; kill=%v", context.Cause(opCtx), killErr) + select { + case lateErr := <-done: + resolved = true + if execErr == nil { + execErr = lateErr + } + _ = ddlConn.Close() + case <-time.After(ddlSettleGrace): + go reapDDLConn(done, ddlConn) + } + } + observeDeadline := maxTime(deadline, time.Now().Add(ddlObservationGrace)) + status, topologyErr := observeStableObjectTopology(db, c, observeDeadline) + if topologyErr != nil { + return topologyUnknown, resolved, fmt.Errorf("rename=%v topology=%w", execErr, topologyErr) + } + return status, resolved, execErr +} + +func waitForExactPendingRename(ctx context.Context, control *sql.Conn, c config, barrierFacts, ddlFacts sessionFacts, tagged string) error { + ticker := time.NewTicker(10 * time.Millisecond) + defer ticker.Stop() + for { + if err := assertExactPendingRename(ctx, control, c, barrierFacts, ddlFacts, tagged); err == nil { + return nil + } + select { + case <-ctx.Done(): + return context.Cause(ctx) + case <-ticker.C: + } + } +} + +func assertExactPendingRename(ctx context.Context, control *sql.Conn, c config, barrierFacts, ddlFacts sessionFacts, tagged string) error { + var database sql.NullString + var user, command string + var info sql.NullString + if err := control.QueryRowContext(ctx, "SELECT DB,USER,COMMAND,INFO FROM information_schema.PROCESSLIST WHERE ID=?", ddlFacts.connectionID).Scan(&database, &user, &command, &info); err != nil { + return err + } + expectedUser := strings.SplitN(ddlFacts.currentUser, "@", 2)[0] + if !database.Valid || database.String != c.schema || user != expectedUser || command != "Query" || !info.Valid || !sameDDL(info.String, tagged) { + return fmt.Errorf("pending process identity mismatch") + } + rows, err := control.QueryContext(ctx, "SELECT COALESCE(t.PROCESSLIST_ID,0),m.LOCK_STATUS FROM performance_schema.metadata_locks m LEFT JOIN performance_schema.threads t ON t.THREAD_ID=m.OWNER_THREAD_ID WHERE m.OBJECT_TYPE='TABLE' AND m.OBJECT_SCHEMA=? AND m.OBJECT_NAME IN (?,?,?)", c.schema, c.source, c.target, c.old) + if err != nil { + return err + } + defer rows.Close() + pendingRename := 0 + for rows.Next() { + var owner int64 + var status string + if err := rows.Scan(&owner, &status); err != nil { + return err + } + switch { + case mdlOwnerAllowed(owner, status, barrierFacts.connectionID, ddlFacts.connectionID): + if owner == ddlFacts.connectionID && status == "PENDING" { + pendingRename++ + } + default: + return fmt.Errorf("unexpected MDL owner=%d status=%s", owner, status) + } + } + if err := rows.Err(); err != nil { + return err + } + if pendingRename < 1 { + return fmt.Errorf("exact RENAME has no pending table MDL") + } + return nil +} + +func mdlOwnerAllowed(owner int64, status string, barrierID, ddlID int64) bool { + return (owner == barrierID && status == "GRANTED") || + (owner == ddlID && (status == "GRANTED" || status == "PENDING")) +} + +func observeStableObjectTopology(db *sql.DB, c config, deadline time.Time) (topologyStatus, error) { + states := make([]topologyStatus, 0, 2) + for i := 0; i < 2; i++ { + observeCtx, cancel := boundedBackground(deadline, 300*time.Millisecond) + conn, _, err := openVerifiedConn(observeCtx, db, c, "fresh-topology-observer") + if err != nil { + cancel() + return topologyUnknown, err + } + var o objectTopology + for name, dst := range map[string]*bool{c.source: &o.source, c.target: &o.target, c.old: &o.old} { + var n int + if err := conn.QueryRowContext(observeCtx, "SELECT COUNT(*) FROM information_schema.tables WHERE table_schema=? AND table_name=?", c.schema, name).Scan(&n); err != nil { + _ = conn.Close() + cancel() + return topologyUnknown, err + } + *dst = n == 1 + } + _ = conn.Close() + cancel() + states = append(states, classifyTopology(o)) + if i == 0 { + time.Sleep(minDuration(25*time.Millisecond, positiveRemaining(deadline))) + } + } + if states[0] != states[1] || states[0] == topologyUnknown { + return topologyUnknown, fmt.Errorf("fresh topology observations unstable: %v", states) + } + return states[0], nil +} + +func recoverCutover(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + deadline := time.Now().Add(maxDuration(c.ddlTimeout, 2*c.statementTimeout)) + status, err := observeStableObjectTopology(db, c, deadline) + if err != nil { + return err + } + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if err := ev.emit("recovery_diagnosis", map[string]any{"topology": status, "phase": state.phase, "generation": state.generation}); err != nil { + return err + } + switch status { + case topologyPreCutover: + if state.phase == "rollback-intent" || state.phase == "rollback-reconcile" { + if state.phase == "rollback-intent" { + if !state.ddlNonce.Valid || !state.ddlOperation.Valid || state.ddlOperation.String != rollbackOperation { + return fmt.Errorf("rollback PRE checkpoint lacks exact persisted RENAME intent") + } + if err := assertNoExactDDLInFlight(ctx, db, c, rollbackOperation, state.ddlNonce.String, rollbackStatement(c)); err != nil { + return err + } + } + return stabilizePreRollback(ctx, db, c, ev) + } + if state.phase == "fresh" { + _, err := assertCutoverPreconditions(ctx, db, c) + return err + } + if state.phase != "ddl-intent" || !state.ddlOperation.Valid || state.ddlOperation.String != cutoverOperation || !state.ddlNonce.Valid { + return fmt.Errorf("PRE topology has unexpected checkpoint phase=%s", state.phase) + } + if err := assertNoExactDDLInFlight(ctx, db, c, cutoverOperation, state.ddlNonce.String, renameStatement(c)); err != nil { + return err + } + if err := assertPreCutoverTriggerTopology(ctx, db, c); err != nil { + return err + } + return clearPreCutoverIntent(ctx, db, c, state) + case topologyPostCutover: + if state.phase == "rollback-intent" { + if !state.ddlNonce.Valid || !state.ddlOperation.Valid || state.ddlOperation.String != rollbackOperation { + return fmt.Errorf("rollback POST checkpoint lacks exact persisted RENAME intent") + } + if err := assertNoExactDDLInFlight(ctx, db, c, rollbackOperation, state.ddlNonce.String, rollbackStatement(c)); err != nil { + return err + } + if err := clearPostCutoverRollbackIntent(ctx, db, c, state); err != nil { + return err + } + state, err = loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + return assertRollbackReady(ctx, db, c, state) + } + return stabilizePostCutover(ctx, db, c, ev) + default: + return fmt.Errorf("recovery refuses UNKNOWN topology") + } +} + +func stabilizePreRollback(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if state.phase != "rollback-intent" && state.phase != "rollback-reconcile" { + return fmt.Errorf("PRE rollback stabilization requires rollback intent, got %s", state.phase) + } + if err := assertOwned(ctx, db, c, c.target); err != nil { + return fmt.Errorf("restored compact target ownership: %w", err) + } + if err := assertFrozenTableFingerprints(ctx, db, c, state); err != nil { + return err + } + if err := ensurePreTriggersAfterRollback(ctx, db, c, state); err != nil { + return err + } + if state.phase == "rollback-intent" { + source, _ := qualified(c.schema, c.source) + var upper sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+source).Scan(&upper); err != nil { + return err + } + cp, _ := qualified(c.schema, c.checkpoint) + result, err := db.ExecContext(ctx, "UPDATE "+cp+" SET phase='rollback-reconcile',final_cutoff_id=?,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='rollback-intent' AND generation=? AND ddl_operation=? AND ddl_nonce=?", upper.Int64, state.generation, rollbackOperation, state.ddlNonce.String) + if err != nil { + return err + } + affected, _ := result.RowsAffected() + if affected != 1 { + return fmt.Errorf("rollback reconcile CAS conflict generation=%d", state.generation) + } + state, err = loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + } + if err := reconcilePreRollbackGap(ctx, db, c, ev); err != nil { + return err + } + state, err = loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if !state.final.Valid || state.last < state.final.Int64 { + return fmt.Errorf("rollback reconcile is incomplete last=%d final=%v", state.last, state.final) + } + if err := removeFilteredRollbackRows(ctx, db, c, ev, state.seed); err != nil { + return err + } + cp, _ := qualified(c.schema, c.checkpoint) + result, err := db.ExecContext(ctx, "UPDATE "+cp+" SET phase='fresh',rollback_base_id=NULL,ddl_operation=NULL,ddl_nonce=NULL,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='rollback-reconcile' AND generation=?", state.generation) + if err != nil { + return err + } + affected, _ := result.RowsAffected() + if affected != 1 { + return fmt.Errorf("rollback completion CAS conflict generation=%d", state.generation) + } + return assertPreCutoverTriggerTopology(ctx, db, c) +} + +func removeFilteredRollbackRows(ctx context.Context, db *sql.DB, c config, ev *evidence, verifiedFloor int64) error { + target, _ := qualified(c.schema, c.target) + var upper sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+target).Scan(&upper); err != nil { + return err + } + for start := verifiedFloor; start < upper.Int64; { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if state.phase != "rollback-reconcile" { + return fmt.Errorf("filtered-row cleanup requires rollback-reconcile checkpoint") + } + if err := assertRuntimeSafe(ctx, db, c, state, true); err != nil { + return err + } + var end sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM (SELECT id FROM "+target+" FORCE INDEX(PRIMARY) WHERE id>? AND id<=? ORDER BY id LIMIT ?) x", start, upper.Int64, c.batchSize).Scan(&end); err != nil { + return err + } + if !end.Valid { + return nil + } + batchCtx, cancel := context.WithTimeout(ctx, c.statementTimeout) + tx, err := db.BeginTx(batchCtx, nil) + if err != nil { + cancel() + return err + } + result, err := tx.ExecContext(batchCtx, "DELETE FROM "+target+" WHERE id>? AND id<=? AND "+filteredPredicate("", c.channelIDs), start, end.Int64) + var removed int64 + if err == nil { + removed, _ = result.RowsAffected() + var remaining int + err = tx.QueryRowContext(batchCtx, "SELECT COUNT(*) FROM "+target+" WHERE id>? AND id<=? AND "+filteredPredicate("", c.channelIDs), start, end.Int64).Scan(&remaining) + if err == nil && remaining != 0 { + err = fmt.Errorf("filtered compact rows remain count=%d range=(%d,%d]", remaining, start, end.Int64) + } + } + if err != nil { + _ = tx.Rollback() + cancel() + return err + } + if err := tx.Commit(); err != nil { + cancel() + return err + } + cancel() + if err := ev.emit("rollback_filtered_cleanup", map[string]any{"start_id": start, "end_id": end.Int64, "removed": removed}); err != nil { + return err + } + start = end.Int64 + } + return nil +} + +func ensurePreTriggersAfterRollback(ctx context.Context, db *sql.DB, c config, state checkpoint) error { + o, t, err := observeTopology(ctx, db, c) + if err != nil { + return err + } + if classifyTopology(o) != topologyPreCutover || t.forward && t.reverse { + return fmt.Errorf("unsafe rollback PRE trigger topology objects=%+v triggers=%+v", o, t) + } + if t.reverse { + spec, err := expectedTriggerSpec(c, "reverse", c.source, state.triggerSQLMode) + if err != nil { + return err + } + spec.table = c.target + name, _ := triggerName("reverse", c.batch) + quoted, _ := quoteIdentifier(name) + if err := runDDL(ctx, db, c, "DROP TRIGGER "+quoted, exactTriggerObserver(spec, false)); err != nil { + return err + } + } + if !t.forward { + query, err := buildForwardTriggerSQL(c) + if err != nil { + return err + } + spec, err := expectedTriggerSpec(c, "forward", c.source, state.triggerSQLMode) + if err != nil { + return err + } + if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { + return err + } + } + for _, event := range []string{"update", "delete"} { + _, current, err := observeTopology(ctx, db, c) + if err != nil { + return err + } + exists, table := current.updateGuard, current.updateGuardTable + if event == "delete" { + exists, table = current.deleteGuard, current.deleteGuardTable + } + if exists && table == c.target { + kind := "guard_" + event + spec, err := expectedTriggerSpec(c, kind, c.target, state.triggerSQLMode) + if err != nil { + return err + } + name, _ := triggerName(kind, c.batch) + quoted, _ := quoteIdentifier(name) + if err := runDDL(ctx, db, c, "DROP TRIGGER "+quoted, exactTriggerObserver(spec, false)); err != nil { + return err + } + exists = false + } + if !exists { + query, err := buildGuardTriggerSQL(c, event, c.source) + if err != nil { + return err + } + spec, err := expectedTriggerSpec(c, "guard_"+event, c.source, state.triggerSQLMode) + if err != nil { + return err + } + if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { + return err + } + } + } + _, final, err := observeTopology(ctx, db, c) + if err != nil || !final.forward || final.reverse || !guardsOnLive(final, c.source) { + return fmt.Errorf("rollback PRE trigger stabilization incomplete: triggers=%+v err=%v", final, err) + } + return nil +} + +func reconcilePreRollbackGap(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + source, _ := qualified(c.schema, c.source) + target, _ := qualified(c.schema, c.target) + cp, _ := qualified(c.schema, c.checkpoint) + copySQL, err := buildCopySQL(c) + if err != nil { + return err + } + for { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if state.phase != "rollback-reconcile" || !state.final.Valid { + return fmt.Errorf("rollback reconcile checkpoint is invalid") + } + if state.last >= state.final.Int64 { + return nil + } + if err := assertRuntimeSafe(ctx, db, c, state, true); err != nil { + return err + } + var end sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM (SELECT id FROM "+source+" FORCE INDEX(PRIMARY) WHERE id>? AND id<=? ORDER BY id LIMIT ?) x", state.last, state.final.Int64, c.batchSize).Scan(&end); err != nil { + return err + } + if !end.Valid { + return nil + } + batchCtx, cancel := context.WithTimeout(ctx, c.statementTimeout) + tx, err := db.BeginTx(batchCtx, nil) + if err != nil { + cancel() + return err + } + if _, err = tx.ExecContext(batchCtx, copySQL, state.last, end.Int64, state.final.Int64); err == nil { + err = verifyWindow(batchCtx, tx, c, source, target, state.last, end.Int64) + } + if err == nil { + var result sql.Result + result, err = tx.ExecContext(batchCtx, "UPDATE "+cp+" SET last_completed_end_id=?,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='rollback-reconcile' AND generation=?", end.Int64, state.generation) + if err == nil { + affected, _ := result.RowsAffected() + if affected != 1 { + err = fmt.Errorf("rollback reconcile checkpoint CAS conflict") + } + } + } + if err != nil { + _ = tx.Rollback() + cancel() + return err + } + if err := tx.Commit(); err != nil { + cancel() + return err + } + cancel() + if err := ev.emit("rollback_reconcile_checkpoint", map[string]any{"end_id": end.Int64, "upper_bound": state.final.Int64}); err != nil { + return err + } + } +} + +func taggedDDL(c config, operation, nonce, statement string) string { + return fmt.Sprintf("/*logs_slim batch=%s operation=%s nonce=%s*/ %s", c.batch, operation, nonce, statement) +} + +func assertNoExactDDLInFlight(ctx context.Context, db *sql.DB, c config, operation, nonce, statement string) error { + conn, facts, err := openVerifiedConn(ctx, db, c, "recover-process-observer") + if err != nil { + return err + } + defer conn.Close() + expectedUser := strings.SplitN(facts.currentUser, "@", 2)[0] + var count int + if err := conn.QueryRowContext(ctx, "SELECT COUNT(*) FROM information_schema.PROCESSLIST WHERE DB=? AND USER=? AND COMMAND='Query' AND INFO=?", c.schema, expectedUser, taggedDDL(c, operation, nonce, statement)).Scan(&count); err != nil { + return fmt.Errorf("prove no exact DDL remains in flight: %w", err) + } + if count != 0 { + return fmt.Errorf("exact DDL is still in flight; persisted intent retained") + } + return nil +} + +func assertPreCutoverTriggerTopology(ctx context.Context, db *sql.DB, c config) error { + o, t, err := observeTopology(ctx, db, c) + if err != nil { + return err + } + if classifyTopology(o) != topologyPreCutover || !t.forward || !t.updateGuard || !t.deleteGuard || t.reverse { + return fmt.Errorf("unsafe PRE trigger topology objects=%+v triggers=%+v", o, t) + } + return assertOwned(ctx, db, c, c.target) +} + +func stabilizePostCutover(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if err := assertOwned(ctx, db, c, c.source); err != nil { + return fmt.Errorf("active POST table ownership: %w", err) + } + if err := assertFrozenTableFingerprints(ctx, db, c, state); err != nil { + return err + } + if state.phase == "rollback-ready" { + return assertRollbackReady(ctx, db, c, state) + } + if state.phase == "ddl-intent" { + if !state.ddlOperation.Valid || state.ddlOperation.String != cutoverOperation || !state.ddlNonce.Valid { + return fmt.Errorf("POST checkpoint lacks exact persisted RENAME intent") + } + old, _ := qualified(c.schema, c.old) + var base sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+old).Scan(&base); err != nil { + return err + } + cp, _ := qualified(c.schema, c.checkpoint) + result, err := db.ExecContext(ctx, "UPDATE "+cp+" SET phase='rollback-gap',rollback_base_id=?,last_completed_end_id=?,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='ddl-intent' AND generation=? AND ddl_operation=? AND ddl_nonce=?", base.Int64, base.Int64, state.generation, cutoverOperation, state.ddlNonce.String) + if err != nil { + return err + } + affected, _ := result.RowsAffected() + if affected != 1 { + return fmt.Errorf("rollback-base CAS conflict generation=%d", state.generation) + } + state, err = loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + } + if state.phase != "rollback-gap" || !state.rollback.Valid { + return fmt.Errorf("POST stabilization requires rollback-gap checkpoint, got %s", state.phase) + } + if err := ensurePostTriggers(ctx, db, c, state); err != nil { + return err + } + if err := reconcileRollbackGap(ctx, db, c, ev); err != nil { + return err + } + if err := fullPostVerify(ctx, db, c); err != nil { + return err + } + state, err = loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + cp, _ := qualified(c.schema, c.checkpoint) + result, err := db.ExecContext(ctx, "UPDATE "+cp+" SET phase='rollback-ready',ddl_operation=NULL,ddl_nonce=NULL,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='rollback-gap' AND generation=?", state.generation) + if err != nil { + return err + } + affected, _ := result.RowsAffected() + if affected != 1 { + return fmt.Errorf("ROLLBACK_READY CAS conflict generation=%d", state.generation) + } + state.phase = "rollback-ready" + state.ddlOperation = sql.NullString{} + state.ddlNonce = sql.NullString{} + return assertRollbackReady(ctx, db, c, state) +} + +func ensurePostTriggers(ctx context.Context, db *sql.DB, c config, state checkpoint) error { + o, t, err := observeTopology(ctx, db, c) + if err != nil { + return err + } + if classifyTopology(o) != topologyPostCutover { + return fmt.Errorf("unsafe POST trigger topology objects=%+v triggers=%+v", o, t) + } + actions, err := postTriggerPlan(t) + if err != nil { + return err + } + if len(actions) > 0 && actions[0] == postTriggerDropForward { + spec, _ := expectedTriggerSpec(c, "forward", c.old, state.triggerSQLMode) + name, _ := triggerName("forward", c.batch) + qn, _ := quoteIdentifier(name) + if err := runDDL(ctx, db, c, "DROP TRIGGER "+qn, exactTriggerObserver(spec, false)); err != nil { + return err + } + } + if len(actions) > 0 { + _, t, err = observeTopology(ctx, db, c) + if err != nil || t.forward || t.reverse { + return fmt.Errorf("forward drop postcondition does not prove zero mirror triggers: triggers=%+v err=%v", t, err) + } + query, err := buildStrictMirrorTriggerSQL(c, c.source, c.old) + if err != nil { + return err + } + spec, err := expectedTriggerSpec(c, "reverse", c.source, state.triggerSQLMode) + if err != nil { + return err + } + if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { + return err + } + } + for _, event := range []string{"update", "delete"} { + if err := ensurePostGuard(ctx, db, c, state, event); err != nil { + return err + } + } + _, final, err := observeTopology(ctx, db, c) + if err != nil || final.forward || !final.reverse || !guardsOnLive(final, c.source) { + return fmt.Errorf("POST trigger stabilization incomplete: triggers=%+v err=%v", final, err) + } + return nil +} + +type postGuardAction string + +const ( + postGuardDropOld postGuardAction = "drop-old" + postGuardCreateLive postGuardAction = "create-live" +) + +func postGuardPlan(exists bool, table, live, old string) ([]postGuardAction, error) { + switch { + case !exists: + return []postGuardAction{postGuardCreateLive}, nil + case table == live: + return nil, nil + case table == old: + return []postGuardAction{postGuardDropOld, postGuardCreateLive}, nil + default: + return nil, fmt.Errorf("guard is attached to unexpected table %q", table) + } +} + +func ensurePostGuard(ctx context.Context, db *sql.DB, c config, state checkpoint, event string) error { + _, t, err := observeTopology(ctx, db, c) + if err != nil { + return err + } + exists, table := t.updateGuard, t.updateGuardTable + if event == "delete" { + exists, table = t.deleteGuard, t.deleteGuardTable + } + actions, err := postGuardPlan(exists, table, c.source, c.old) + if err != nil { + return err + } + kind := "guard_" + event + for _, action := range actions { + switch action { + case postGuardDropOld: + spec, err := expectedTriggerSpec(c, kind, c.old, state.triggerSQLMode) + if err != nil { + return err + } + name, _ := triggerName(kind, c.batch) + quoted, _ := quoteIdentifier(name) + if err := runDDL(ctx, db, c, "DROP TRIGGER "+quoted, exactTriggerObserver(spec, false)); err != nil { + return err + } + case postGuardCreateLive: + query, err := buildGuardTriggerSQL(c, event, c.source) + if err != nil { + return err + } + spec, err := expectedTriggerSpec(c, kind, c.source, state.triggerSQLMode) + if err != nil { + return err + } + if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { + return err + } + } + } + return nil +} + +func guardsOnLive(t triggerTopology, live string) bool { + return t.updateGuard && t.deleteGuard && t.updateGuardTable == live && t.deleteGuardTable == live +} + +type postTriggerAction string + +const ( + postTriggerDropForward postTriggerAction = "drop-forward" + postTriggerCreateReverse postTriggerAction = "create-reverse" +) + +func postTriggerPlan(t triggerTopology) ([]postTriggerAction, error) { + if t.forward && t.reverse { + return nil, fmt.Errorf("forward and reverse triggers must never coexist") + } + if t.reverse { + return nil, nil + } + if t.forward { + return []postTriggerAction{postTriggerDropForward, postTriggerCreateReverse}, nil + } + return []postTriggerAction{postTriggerCreateReverse}, nil +} + +func reconcileRollbackGap(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + source, _ := qualified(c.schema, c.source) + old, _ := qualified(c.schema, c.old) + cp, _ := qualified(c.schema, c.checkpoint) + columns := strings.Join(logColumns, ", ") + copySQL := "INSERT INTO " + old + " (" + columns + ") SELECT " + columns + " FROM " + source + " WHERE id>? AND id<=? ON DUPLICATE KEY UPDATE id=" + old + ".id" + var upper sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+source).Scan(&upper); err != nil { + return err + } + for { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if state.last >= upper.Int64 { + return nil + } + var end sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM (SELECT id FROM "+source+" WHERE id>? AND id<=? ORDER BY id LIMIT ?) x", state.last, upper.Int64, c.batchSize).Scan(&end); err != nil { + return err + } + if !end.Valid { + return nil + } + batchCtx, cancel := context.WithTimeout(ctx, c.statementTimeout) + tx, err := db.BeginTx(batchCtx, nil) + if err != nil { + cancel() + return err + } + if _, err = tx.ExecContext(batchCtx, copySQL, state.last, end.Int64); err == nil { + err = verifyMirrorWindow(batchCtx, tx, source, old, state.last, end.Int64) + } + if err == nil { + var result sql.Result + result, err = tx.ExecContext(batchCtx, "UPDATE "+cp+" SET last_completed_end_id=?,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND phase='rollback-gap' AND generation=?", end.Int64, state.generation) + if err == nil { + affected, _ := result.RowsAffected() + if affected != 1 { + err = fmt.Errorf("rollback-gap checkpoint CAS conflict") + } + } + } + if err != nil { + _ = tx.Rollback() + cancel() + return err + } + if err := tx.Commit(); err != nil { + cancel() + return err + } + cancel() + if err := ev.emit("rollback_gap_checkpoint", map[string]any{"end_id": end.Int64, "upper_bound": upper.Int64}); err != nil { + return err + } + } +} + +func verifyMirrorWindow(ctx context.Context, q queryRower, source, old string, start, end int64) error { + var count int + query := "SELECT COUNT(*) FROM " + source + " s LEFT JOIN " + old + " o ON o.id=s.id WHERE s.id>? AND s.id<=? AND (o.id IS NULL OR NOT (" + rowEqualitySQL("s", "o") + "))" + if err := q.QueryRowContext(ctx, query, start, end).Scan(&count); err != nil { + return err + } + if count != 0 { + return fmt.Errorf("rollback mirror verification failed count=%d range=(%d,%d]", count, start, end) + } + return nil +} + +func fullPostVerify(ctx context.Context, db *sql.DB, c config) error { + source, _ := qualified(c.schema, c.source) + old, _ := qualified(c.schema, c.old) + var oldUpper, sourceUpper sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+old).Scan(&oldUpper); err != nil { + return err + } + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+source).Scan(&sourceUpper); err != nil { + return err + } + upper := maxInt64(oldUpper.Int64, sourceUpper.Int64) + for start := int64(0); start < upper; { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if err := assertRuntimeSafe(ctx, db, c, state, false); err != nil { + return err + } + var oldEnd, sourceEnd sql.NullInt64 + endSQL := "SELECT MAX(id) FROM (SELECT id FROM %s FORCE INDEX(PRIMARY) WHERE id>? AND id<=? ORDER BY id LIMIT ?) x" + if err := db.QueryRowContext(ctx, fmt.Sprintf(endSQL, old), start, upper, c.batchSize).Scan(&oldEnd); err != nil { + return err + } + if err := db.QueryRowContext(ctx, fmt.Sprintf(endSQL, source), start, upper, c.batchSize).Scan(&sourceEnd); err != nil { + return err + } + end := nextPostVerifyEnd(oldEnd, sourceEnd) + if !end.Valid { + break + } + batchCtx, cancel := context.WithTimeout(ctx, c.statementTimeout) + var oldMissing, oldMismatch, sourceMissing, sourceMismatch int + query := "SELECT COUNT(*) FROM " + old + " o LEFT JOIN " + source + " s ON s.id=o.id WHERE o.id>? AND o.id<=? AND " + retainedPredicate("o", c.channelIDs) + " AND s.id IS NULL" + err = db.QueryRowContext(batchCtx, query, start, end.Int64).Scan(&oldMissing) + if err == nil { + query = "SELECT COUNT(*) FROM " + old + " o JOIN " + source + " s ON s.id=o.id WHERE o.id>? AND o.id<=? AND " + retainedPredicate("o", c.channelIDs) + " AND NOT (" + rowEqualitySQL("o", "s") + ")" + err = db.QueryRowContext(batchCtx, query, start, end.Int64).Scan(&oldMismatch) + } + if err == nil { + query = "SELECT COUNT(*) FROM " + source + " s LEFT JOIN " + old + " o ON o.id=s.id WHERE s.id>? AND s.id<=? AND o.id IS NULL" + err = db.QueryRowContext(batchCtx, query, start, end.Int64).Scan(&sourceMissing) + } + if err == nil { + query = "SELECT COUNT(*) FROM " + source + " s JOIN " + old + " o ON o.id=s.id WHERE s.id>? AND s.id<=? AND NOT (" + rowEqualitySQL("s", "o") + ")" + err = db.QueryRowContext(batchCtx, query, start, end.Int64).Scan(&sourceMismatch) + } + cancel() + if err != nil || oldMissing != 0 || oldMismatch != 0 || sourceMissing != 0 || sourceMismatch != 0 { + return fmt.Errorf("full POST verification failed range=(%d,%d] old_missing=%d old_mismatch=%d source_missing=%d source_mismatch=%d err=%v", start, end.Int64, oldMissing, oldMismatch, sourceMissing, sourceMismatch, err) + } + start = end.Int64 + if c.batchDelay > 0 { + select { + case <-ctx.Done(): + return context.Cause(ctx) + case <-time.After(c.batchDelay): + } + } + } + return nil +} + +func nextPostVerifyEnd(oldEnd, sourceEnd sql.NullInt64) sql.NullInt64 { + switch { + case oldEnd.Valid && sourceEnd.Valid && oldEnd.Int64 <= sourceEnd.Int64: + return oldEnd + case oldEnd.Valid && sourceEnd.Valid: + return sourceEnd + case oldEnd.Valid: + return oldEnd + default: + return sourceEnd + } +} + +func maxInt64(a, b int64) int64 { + if a > b { + return a + } + return b +} + +func assertRollbackReady(ctx context.Context, db *sql.DB, c config, state checkpoint) error { + if state.phase != "rollback-ready" || !state.rollback.Valid || state.ddlOperation.Valid || state.ddlNonce.Valid { + return fmt.Errorf("checkpoint is not ROLLBACK_READY") + } + o, t, err := observeTopology(ctx, db, c) + if err != nil { + return err + } + if classifyTopology(o) != topologyPostCutover || t.forward || !t.reverse || !guardsOnLive(t, c.source) { + return fmt.Errorf("ROLLBACK_READY topology mismatch objects=%+v triggers=%+v", o, t) + } + return nil +} + +func maxDuration(a, b time.Duration) time.Duration { + if a > b { + return a + } + return b +} + +func safeNext(max uint64) (uint64, error) { + if max == math.MaxUint64 { + return 0, fmt.Errorf("AUTO_INCREMENT exhausted") + } + return max + 1, nil +} diff --git a/cmd/logs_slimming/ddl.go b/cmd/logs_slimming/ddl.go new file mode 100644 index 00000000000..b7d59171170 --- /dev/null +++ b/cmd/logs_slimming/ddl.go @@ -0,0 +1,229 @@ +package main + +import ( + "context" + "crypto/rand" + "crypto/sha256" + "database/sql" + "encoding/hex" + "fmt" + "strings" + "time" +) + +type ddlState string + +const ( + ddlPre ddlState = "PRE" + ddlPost ddlState = "POST" + ddlUnknown ddlState = "UNKNOWN" +) + +type ddlObserver func(context.Context, *sql.Conn, string) (ddlState, error) + +func sameDDL(expected, actual string) bool { + return strings.Join(strings.Fields(expected), " ") == strings.Join(strings.Fields(actual), " ") +} + +func ddlOperation(statement string) string { + fields := strings.Fields(strings.ToUpper(statement)) + if len(fields) == 0 { + return "EMPTY" + } + if len(fields) == 1 { + return fields[0] + } + return fields[0] + "_" + fields[1] +} + +func ddlNonce() (string, error) { + data := make([]byte, 12) + if _, err := rand.Read(data); err != nil { + return "", err + } + return hex.EncodeToString(data), nil +} + +// runDDL bounds the complete operation, including kill and two independent +// postcondition observations. If KILL does not resolve the driver call within +// the bounded grace period, the physical DDL connection is closed only by a +// background reaper after Exec returns; the caller never blocks in Close. +func runDDL(ctx context.Context, db *sql.DB, c config, statement string, observe ddlObserver) error { + started := time.Now() + deadline := started.Add(c.ddlTimeout) + opCtx, cancelOperation := context.WithDeadline(ctx, deadline) + defer cancelOperation() + operation := ddlOperation(statement) + nonce, err := ddlNonce() + if err != nil { + return fmt.Errorf("generate DDL nonce: %w", err) + } + ddlConn, ddlFacts, err := openVerifiedConn(opCtx, db, c, "ddl") + if err != nil { + return err + } + control, controlFacts, err := openVerifiedConn(opCtx, db, c, "control") + if err != nil { + _ = ddlConn.Close() + return err + } + defer control.Close() + if controlFacts.connectionID == ddlFacts.connectionID { + _ = ddlConn.Close() + return fmt.Errorf("DDL and control connections unexpectedly share id %d", ddlFacts.connectionID) + } + if _, err = ddlConn.ExecContext(opCtx, fmt.Sprintf("SET SESSION lock_wait_timeout = %d", c.lockWaitSeconds)); err != nil { + _ = ddlConn.Close() + return err + } + tagged := fmt.Sprintf("/*logs_slim batch=%s operation=%s nonce=%s*/ %s", c.batch, operation, nonce, statement) + ev, ok := ctx.Value(evidenceContextKey{}).(*evidence) + if !ok { + _ = ddlConn.Close() + return fmt.Errorf("DDL requires fail-closed evidence context") + } + statementHash := fmt.Sprintf("%x", sha256.Sum256([]byte(statement))) + if err := ev.emit("ddl_intent", map[string]any{"batch": c.batch, "operation": operation, "nonce": nonce, "connection_id": ddlFacts.connectionID, "statement_sha256": statementHash}); err != nil { + _ = ddlConn.Close() + return fmt.Errorf("write DDL intent evidence: %w", err) + } + + done := make(chan error, 1) + go func() { + _, execErr := ddlConn.ExecContext(opCtx, tagged) + done <- execErr + }() + watchdogAt := time.Until(deadline) - 600*time.Millisecond + if watchdogAt < 50*time.Millisecond { + watchdogAt = 50 * time.Millisecond + } + timer := time.NewTimer(watchdogAt) + defer timer.Stop() + resolved := false + var execErr error + var killErr error + select { + case execErr = <-done: + resolved = true + case <-timer.C: + killCtx, cancel := boundedBackground(deadline, 250*time.Millisecond) + killErr = killExactDDL(killCtx, control, c, ctx, ddlFacts, tagged) + cancel() + // Cancellation is mandatory even when the exact KILL could not be proven. + // A failed KILL is an UNKNOWN outcome, never permission to return early + // while the server-side statement may still complete later. + cancelOperation() + grace := time.NewTimer(ddlSettleGrace) + select { + case execErr = <-done: + resolved = true + case <-grace.C: + } + grace.Stop() + } + if resolved { + _ = ddlConn.Close() + } else { + go reapDDLConn(done, ddlConn) + } + + observeDeadline := maxTime(deadline, time.Now().Add(ddlObservationGrace)) + states := make([]ddlState, 0, 2) + for i := 0; i < 2; i++ { + observeCtx, cancel := boundedBackground(observeDeadline, 250*time.Millisecond) + observer, _, openErr := openVerifiedConn(observeCtx, db, c, "observer") + if openErr != nil { + cancel() + return openErr + } + state, observeErr := observe(observeCtx, observer, ddlFacts.sqlMode) + _ = observer.Close() + cancel() + if observeErr != nil { + return observeErr + } + states = append(states, state) + if i == 0 { + time.Sleep(minDuration(25*time.Millisecond, positiveRemaining(deadline))) + } + } + if !resolved || killErr != nil { + if err := ev.emit("ddl_postcondition", map[string]any{"operation": operation, "nonce": nonce, "statement_sha256": statementHash, "state": fmt.Sprint(states), "result": "unknown", "exec_resolved": resolved, "kill_error": killErr}); err != nil { + return fmt.Errorf("write unresolved DDL postcondition evidence: %w", err) + } + return fmt.Errorf("DDL outcome UNKNOWN: exec_resolved=%t kill=%v states=%v", resolved, killErr, states) + } + if states[0] != states[1] || states[0] == ddlUnknown { + if err := ev.emit("ddl_postcondition", map[string]any{"operation": operation, "nonce": nonce, "statement_sha256": statementHash, "state": fmt.Sprint(states), "result": "unknown"}); err != nil { + return fmt.Errorf("write unknown DDL postcondition evidence: %w", err) + } + return fmt.Errorf("DDL postcondition is unstable or unknown: %v", states) + } + if states[0] == ddlPost { + if err := ev.emit("ddl_postcondition", map[string]any{"operation": operation, "nonce": nonce, "statement_sha256": statementHash, "state": string(ddlPost), "elapsed_ms": time.Since(started).Milliseconds()}); err != nil { + return fmt.Errorf("write DDL postcondition evidence: %w", err) + } + return nil + } + if execErr != nil { + return execErr + } + return fmt.Errorf("DDL returned but postcondition remained PRE") +} + +const ( + ddlSettleGrace = 500 * time.Millisecond + ddlObservationGrace = 750 * time.Millisecond +) + +func maxTime(a, b time.Time) time.Time { + if a.After(b) { + return a + } + return b +} + +func killExactDDL(ctx context.Context, control *sql.Conn, c config, lockCtx context.Context, ddlFacts sessionFacts, tagged string) error { + if _, err := verifyPhysicalConn(ctx, control, c, "control-watchdog"); err != nil { + return err + } + var database sql.NullString + var user, command string + var info sql.NullString + if err := control.QueryRowContext(ctx, "SELECT DB,USER,COMMAND,INFO FROM information_schema.PROCESSLIST WHERE ID=?", ddlFacts.connectionID).Scan(&database, &user, &command, &info); err != nil { + return fmt.Errorf("watchdog cannot prove DDL connection identity: %w", err) + } + lockOwner, ok := lockCtx.Value(lockOwnerContextKey{}).(int64) + var liveOwner sql.NullInt64 + lockName := advisoryLockName(c) + lockErr := control.QueryRowContext(ctx, "SELECT IS_USED_LOCK(?)", lockName).Scan(&liveOwner) + expectedUser := strings.SplitN(ddlFacts.currentUser, "@", 2)[0] + if !ok || lockErr != nil || !liveOwner.Valid || liveOwner.Int64 != lockOwner || !database.Valid || database.String != c.schema || user != expectedUser || command != "Query" || !info.Valid || !sameDDL(tagged, info.String) { + return fmt.Errorf("watchdog refused KILL: connection %d does not match exact batch/operation/nonce/DB/USER/lock owner", ddlFacts.connectionID) + } + if _, err := control.ExecContext(ctx, fmt.Sprintf("KILL QUERY %d", ddlFacts.connectionID)); err != nil { + return fmt.Errorf("watchdog KILL QUERY %d: %w", ddlFacts.connectionID, err) + } + return nil +} + +func reapDDLConn(done <-chan error, conn *sql.Conn) { + <-done + _ = conn.Close() +} + +func boundedBackground(deadline time.Time, maximum time.Duration) (context.Context, context.CancelFunc) { + remaining := positiveRemaining(deadline) + if remaining > maximum { + remaining = maximum + } + return context.WithTimeout(context.Background(), remaining) +} + +func positiveRemaining(deadline time.Time) time.Duration { + remaining := time.Until(deadline) + if remaining <= 0 { + return time.Nanosecond + } + return remaining +} diff --git a/cmd/logs_slimming/evidence.go b/cmd/logs_slimming/evidence.go new file mode 100644 index 00000000000..5015fddbf9e --- /dev/null +++ b/cmd/logs_slimming/evidence.go @@ -0,0 +1,89 @@ +package main + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "strings" + "sync" + "time" +) + +type evidenceSink interface { + Write([]byte) (int, error) +} + +type evidence struct { + mu sync.Mutex + sink evidenceSink + secrets []string +} + +func newEvidence(sink evidenceSink, secrets []string) *evidence { + filtered := make([]string, 0, len(secrets)) + for _, secret := range secrets { + if secret != "" { + filtered = append(filtered, secret) + } + } + return &evidence{sink: sink, secrets: filtered} +} + +func (e *evidence) emit(event string, fields map[string]any) error { + record := make(map[string]any, len(fields)+2) + record["timestamp"] = time.Now().UTC().Format(time.RFC3339Nano) + record["event"] = event + for key, value := range fields { + record[key] = e.redactValue(value) + } + var buffer bytes.Buffer + encoder := json.NewEncoder(&buffer) + encoder.SetEscapeHTML(false) + err := encoder.Encode(record) + data := bytes.TrimSuffix(buffer.Bytes(), []byte{'\n'}) + if err != nil { + data = []byte(fmt.Sprintf(`{"timestamp":%q,"event":"evidence_marshal_failed","error":%q}`, time.Now().UTC().Format(time.RFC3339Nano), e.redactString(err.Error()))) + } + e.mu.Lock() + defer e.mu.Unlock() + _, err = e.sink.Write(append(data, '\n')) + return err +} + +func (e *evidence) redactValue(value any) any { + switch typed := value.(type) { + case error: + return e.redactString(typed.Error()) + case string: + return e.redactString(typed) + default: + return typed + } +} + +func (e *evidence) redactString(value string) string { + for _, secret := range e.secrets { + value = strings.ReplaceAll(value, secret, "") + } + return value +} + +type memoryEvidence struct { + mu sync.Mutex + b strings.Builder +} + +func (m *memoryEvidence) Write(data []byte) (int, error) { + m.mu.Lock() + defer m.mu.Unlock() + return m.b.Write(data) +} + +func (m *memoryEvidence) String() string { + m.mu.Lock() + defer m.mu.Unlock() + return m.b.String() +} + +var _ io.Writer = (*memoryEvidence)(nil) diff --git a/cmd/logs_slimming/main.go b/cmd/logs_slimming/main.go new file mode 100644 index 00000000000..21200f84a96 --- /dev/null +++ b/cmd/logs_slimming/main.go @@ -0,0 +1,101 @@ +package main + +import ( + "context" + "database/sql" + "fmt" + "io" + "os" + "os/signal" + "syscall" + "time" + + _ "github.com/go-sql-driver/mysql" +) + +func main() { + cfg, err := parseConfig(os.Args[1:]) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(2) + } + var sink io.Writer = os.Stdout + var file *os.File + if cfg.evidencePath != "" { + file, err = os.OpenFile(cfg.evidencePath, os.O_CREATE|os.O_APPEND|os.O_WRONLY, 0o600) + if err != nil { + fmt.Fprintln(os.Stderr, err) + os.Exit(1) + } + defer file.Close() + sink = file + } + ev := newEvidence(sink, []string{cfg.dsn}) + if cfg.dsn == "" { + if err := ev.emit("refused", map[string]any{"error": dsnEnvironment + " is required"}); err != nil { + fmt.Fprintln(os.Stderr, "write evidence:", err) + } + os.Exit(2) + } + db, err := sql.Open("mysql", cfg.dsn) + if err != nil { + if emitErr := ev.emit("failed", map[string]any{"error": err}); emitErr != nil { + fmt.Fprintln(os.Stderr, "write evidence:", emitErr) + } + os.Exit(1) + } + defer db.Close() + db.SetMaxOpenConns(8) + db.SetMaxIdleConns(0) + db.SetConnMaxLifetime(5 * time.Minute) + ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGINT, syscall.SIGTERM) + defer stop() + if err := runCommand(ctx, db, cfg, ev); err != nil { + if emitErr := ev.emit("failed", map[string]any{"command": cfg.command, "error": err}); emitErr != nil { + fmt.Fprintln(os.Stderr, "write failure evidence:", emitErr) + } + fmt.Fprintln(os.Stderr, "logs_slimming:", ev.redactString(err.Error())) + os.Exit(1) + } + if err := ev.emit("completed", map[string]any{"command": cfg.command, "schema": cfg.schema, "batch": cfg.batch}); err != nil { + fmt.Fprintln(os.Stderr, "write completion evidence:", err) + os.Exit(1) + } +} + +func runCommand(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + if err := preflight(ctx, db, c, ev); err != nil { + return err + } + if c.command != "preflight" && c.command != "prepare" { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if err := assertFrozenTableFingerprints(ctx, db, c, state); err != nil { + return err + } + } + switch c.command { + case "preflight": + return nil + case "prepare": + return withLock(ctx, db, c, ev, func(lockCtx context.Context) error { return prepare(lockCtx, db, c) }) + case "backfill", "reconcile": + return withLock(ctx, db, c, ev, func(lockCtx context.Context) error { return backfill(lockCtx, db, c, ev) }) + case "install-forward-trigger": + return withLock(ctx, db, c, ev, func(lockCtx context.Context) error { return installForward(lockCtx, db, c) }) + case "verify": + return withLock(ctx, db, c, ev, func(lockCtx context.Context) error { return verify(lockCtx, db, c, ev) }) + case "recover": + return withLock(ctx, db, c, ev, func(lockCtx context.Context) error { return recover(lockCtx, db, c, ev) }) + case "cleanup": + return withLock(ctx, db, c, ev, func(lockCtx context.Context) error { return cleanup(lockCtx, db, c) }) + case "cutover": + return withLock(ctx, db, c, ev, func(lockCtx context.Context) error { return cutover(lockCtx, db, c, ev) }) + case "rollback": + return withLock(ctx, db, c, ev, func(lockCtx context.Context) error { return rollback(lockCtx, db, c, ev) }) + default: + return fmt.Errorf("unsupported command %q", c.command) + } +} diff --git a/cmd/logs_slimming/ops.go b/cmd/logs_slimming/ops.go new file mode 100644 index 00000000000..4b5bf70f460 --- /dev/null +++ b/cmd/logs_slimming/ops.go @@ -0,0 +1,1083 @@ +package main + +import ( + "context" + "crypto/sha256" + "database/sql" + "encoding/base64" + "errors" + "fmt" + "os" + "slices" + "strings" + "time" + + "github.com/go-sql-driver/mysql" +) + +type lockOwnerContextKey struct{} +type evidenceContextKey struct{} + +func ownershipMarker(c config) string { + payload := strings.Join([]string{c.expectedProject, c.expectedInstance, c.expectedServerUUID, c.schema, c.source, c.target, c.old, c.checkpoint, c.batch, channelList(c.channelIDs), auditedSchemaContractHash()}, "|") + return fmt.Sprintf("logs-slimming:%s:%x", c.batch, sha256.Sum256([]byte(payload))) +} + +func advisoryLockName(c config) string { + payload := strings.Join([]string{c.expectedProject, c.expectedInstance, c.schema, c.source}, "|") + digest := sha256.Sum256([]byte(payload)) + // MySQL user-level lock names are limited to 64 characters. + return "logs-slimming:" + base64.RawURLEncoding.EncodeToString(digest[:]) +} + +func preflight(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + if os.Getenv("GOOGLE_CLOUD_PROJECT") != c.expectedProject { + return fmt.Errorf("GOOGLE_CLOUD_PROJECT does not equal expected-project") + } + if os.Getenv("CLOUD_SQL_INSTANCE") != c.expectedInstance { + return fmt.Errorf("CLOUD_SQL_INSTANCE does not equal expected-instance") + } + var schema, hostname, uuid, version, isolation, binlog, currentUser string + var autoMode int + err := db.QueryRowContext(ctx, "SELECT DATABASE(), @@hostname, @@server_uuid, VERSION(), @@transaction_isolation, @@binlog_format, @@innodb_autoinc_lock_mode,CURRENT_USER()").Scan(&schema, &hostname, &uuid, &version, &isolation, &binlog, &autoMode, ¤tUser) + if err != nil { + return fmt.Errorf("identity query: %w", err) + } + if schema != c.schema || hostname != c.expectedHostname || uuid != c.expectedServerUUID || currentUser != c.expectedDatabaseUser { + return fmt.Errorf("database identity mismatch schema=%q hostname=%q server_uuid=%q current_user=%q", schema, hostname, uuid, currentUser) + } + if !strings.HasPrefix(version, "8.0.") || isolation != "READ-COMMITTED" || binlog != "ROW" || autoMode != 2 { + return fmt.Errorf("unsupported session facts version=%q isolation=%q binlog=%q autoinc_mode=%d", version, isolation, binlog, autoMode) + } + if err := assertCutoverObservationAccess(ctx, db); err != nil { + return err + } + var sourceExists int + if err := db.QueryRowContext(ctx, "SELECT COUNT(*) FROM information_schema.tables WHERE table_schema=? AND table_name=?", c.schema, c.source).Scan(&sourceExists); err != nil { + return err + } + if sourceExists != 1 { + return fmt.Errorf("source table is missing") + } + sourceFingerprint, err := schemaFingerprint(ctx, db, c.schema, c.source) + if err != nil { + return fmt.Errorf("fingerprint source table: %w", err) + } + rows, err := db.QueryContext(ctx, "SELECT id FROM channels WHERE type=57 ORDER BY id") + if err != nil { + return fmt.Errorf("read Codex channel snapshot: %w", err) + } + var currentChannelIDs []int64 + for rows.Next() { + var id int64 + if err := rows.Scan(&id); err != nil { + rows.Close() + return fmt.Errorf("scan Codex channel snapshot: %w", err) + } + currentChannelIDs = append(currentChannelIDs, id) + } + if err := rows.Close(); err != nil { + return fmt.Errorf("close Codex channel snapshot: %w", err) + } + if err := validateChannelSnapshot(c.channelIDs, currentChannelIDs); err != nil { + return err + } + return ev.emit("preflight_passed", map[string]any{"schema": schema, "hostname": hostname, "server_uuid": uuid, "version": version, "project": c.expectedProject, "instance": c.expectedInstance, "batch": c.batch, "channel_ids": c.channelIDs, "source_fingerprint": sourceFingerprint}) +} + +func assertCutoverObservationAccess(ctx context.Context, q queryRower) error { + return probeCutoverObservationAccess(ctx, func(ctx context.Context, table string) error { + var count int + return q.QueryRowContext(ctx, cutoverObservationProbeSQL(table)).Scan(&count) + }) +} + +func probeCutoverObservationAccess(ctx context.Context, probe func(context.Context, string) error) error { + var denied []string + var failed []string + for _, table := range cutoverObservationTables { + if err := probe(ctx, table); err != nil { + var mysqlErr *mysql.MySQLError + if errors.As(err, &mysqlErr) && mysqlErr.Number == 1142 { + denied = append(denied, "performance_schema."+table) + continue + } + failed = append(failed, fmt.Sprintf("performance_schema.%s: %v", table, err)) + } + } + if len(denied) == 0 && len(failed) == 0 { + return nil + } + parts := make([]string, 0, 2) + if len(denied) != 0 { + parts = append(parts, "SELECT denied on "+strings.Join(denied, ", ")) + } + if len(failed) != 0 { + parts = append(parts, "probe execution failed: "+strings.Join(failed, "; ")) + } + return fmt.Errorf("cutover observation unavailable: %s", strings.Join(parts, "; ")) +} + +var cutoverObservationTables = []string{"metadata_locks", "threads"} + +func cutoverObservationProbeSQL(table string) string { + return "SELECT COUNT(*) FROM performance_schema." + table + " WHERE 1=0" +} + +func validateChannelSnapshot(configured, current []int64) error { + if !slices.Equal(configured, current) { + return fmt.Errorf("frozen Codex channel IDs do not match channels.type=57: configured=%v current=%v", configured, current) + } + return nil +} + +func withLock(ctx context.Context, db *sql.DB, c config, ev *evidence, fn func(context.Context) error) error { + conn, facts, err := openVerifiedConn(ctx, db, c, "lock") + if err != nil { + return err + } + defer conn.Close() + lockName := advisoryLockName(c) + var got int + if err := conn.QueryRowContext(ctx, "SELECT GET_LOCK(?,0)", lockName).Scan(&got); err != nil { + return fmt.Errorf("acquire advisory lock: %w", err) + } + if got != 1 { + return fmt.Errorf("advisory lock unavailable: got=%d", got) + } + ownerID := facts.connectionID + lockCtx, cancel := context.WithCancelCause(ctx) + defer cancel(nil) + done := make(chan struct{}) + go func() { + defer close(done) + ticker := time.NewTicker(2 * time.Second) + defer ticker.Stop() + for { + select { + case <-lockCtx.Done(): + return + case <-ticker.C: + var current sql.NullInt64 + checkCtx, stop := context.WithTimeout(lockCtx, time.Second) + _, identityErr := verifyPhysicalConn(checkCtx, conn, c, "lock-heartbeat") + err := conn.QueryRowContext(checkCtx, "SELECT IS_USED_LOCK(?)", lockName).Scan(¤t) + stop() + if identityErr != nil || err != nil || !current.Valid || current.Int64 != ownerID { + cancel(fmt.Errorf("advisory lock lost: owner=%d current=%v identity=%v err=%w", ownerID, current, identityErr, err)) + return + } + } + } + }() + if err := ev.emit("advisory_lock_acquired", map[string]any{"lock": lockName, "connection_id": ownerID}); err != nil { + cancel(err) + <-done + return err + } + lockCtx = context.WithValue(lockCtx, lockOwnerContextKey{}, ownerID) + lockCtx = context.WithValue(lockCtx, evidenceContextKey{}, ev) + err = fn(lockCtx) + lockCause := context.Cause(lockCtx) + cancel(nil) + <-done + if err == nil && lockCause == nil { + var finalOwner sql.NullInt64 + checkCtx, checkStop := context.WithTimeout(context.Background(), time.Second) + checkErr := conn.QueryRowContext(checkCtx, "SELECT IS_USED_LOCK(?)", lockName).Scan(&finalOwner) + checkStop() + if checkErr != nil || !finalOwner.Valid || finalOwner.Int64 != ownerID { + err = fmt.Errorf("advisory lock ownership lost before command completion: %w", checkErr) + } + } + var released sql.NullInt64 + releaseCtx, stop := context.WithTimeout(context.Background(), 2*time.Second) + defer stop() + _ = conn.QueryRowContext(releaseCtx, "SELECT RELEASE_LOCK(?)", lockName).Scan(&released) + if err == nil && lockCause != nil { + err = lockCause + } + return err +} + +func prepare(ctx context.Context, db *sql.DB, c config) error { + marker := ownershipMarker(c) + target, _ := qualified(c.schema, c.target) + source, _ := qualified(c.schema, c.source) + cp, _ := qualified(c.schema, c.checkpoint) + targetExists, targetComment, err := tableIdentity(ctx, db, c.schema, c.target) + if err != nil { + return err + } + if targetExists { + if targetComment != marker { + return fmt.Errorf("refusing to adopt existing unowned target %s", c.target) + } + sourceHash, err := schemaFingerprint(ctx, db, c.schema, c.source) + if err != nil { + return err + } + targetHash, err := schemaFingerprint(ctx, db, c.schema, c.target) + if err != nil { + return err + } + if sourceHash != targetHash { + return fmt.Errorf("owned target schema fingerprint differs from source") + } + } else { + if err := runDDL(ctx, db, c, "CREATE TABLE "+target+" LIKE "+source, tableStateObserver(c.schema, c.target, true)); err != nil { + return err + } + if err := runDDL(ctx, db, c, "ALTER TABLE "+target+" COMMENT='"+marker+"'", tableCommentObserver(c.schema, c.target, marker)); err != nil { + return err + } + } + sourceHash, err := auditedLogsSchemaHash(ctx, db, c.schema, c.source) + if err != nil { + return err + } + sourceFingerprint, err := schemaFingerprint(ctx, db, c.schema, c.source) + if err != nil { + return err + } + targetFingerprint, err := schemaFingerprint(ctx, db, c.schema, c.target) + if err != nil { + return err + } + if targetFingerprint != sourceFingerprint { + return fmt.Errorf("target full schema/index fingerprint differs from source") + } + var sqlMode string + if err := db.QueryRowContext(ctx, "SELECT @@SESSION.sql_mode").Scan(&sqlMode); err != nil { + return err + } + // Enforce append-only semantics before the first checkpoint can authorize a + // backfill. This avoids privileged performance_schema access and prevents an + // UPDATE or DELETE from racing the seed copy. + for _, event := range []string{"update", "delete"} { + query, err := buildGuardTriggerSQL(c, event, c.source) + if err != nil { + return err + } + spec, err := expectedTriggerSpec(c, "guard_"+event, c.source, sqlMode) + if err != nil { + return err + } + if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { + return err + } + } + statement := "CREATE TABLE " + cp + " (id TINYINT NOT NULL PRIMARY KEY, marker VARCHAR(191) NOT NULL, source_schema_sha256 CHAR(64) NOT NULL, source_fingerprint_sha256 CHAR(64) NOT NULL, phase VARCHAR(32) NOT NULL, last_completed_end_id BIGINT NOT NULL, seed_cutoff_id BIGINT NOT NULL, final_cutoff_id BIGINT NULL, rollback_base_id BIGINT NULL, generation BIGINT UNSIGNED NOT NULL, trigger_sql_mode TEXT NOT NULL, ddl_operation VARCHAR(64) NULL, ddl_nonce VARCHAR(64) NULL, baseline_updates BIGINT UNSIGNED NOT NULL, baseline_deletes BIGINT UNSIGNED NOT NULL, updated_at DATETIME(6) NOT NULL) ENGINE=InnoDB COMMENT='" + marker + "'" + cpExists, cpComment, err := tableIdentity(ctx, db, c.schema, c.checkpoint) + if err != nil { + return err + } + if cpExists { + if cpComment != marker { + return fmt.Errorf("refusing to adopt existing unowned checkpoint %s", c.checkpoint) + } + if err := assertCheckpointSchema(ctx, db, c); err != nil { + return err + } + state, err := loadCheckpoint(ctx, db, c) + if err == nil { + return assertFrozenTableFingerprints(ctx, db, c, state) + } + if err != sql.ErrNoRows { + return err + } + } else { + if err := runDDL(ctx, db, c, statement, tableCommentObserver(c.schema, c.checkpoint, marker)); err != nil { + return err + } + } + var maxID sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+source).Scan(&maxID); err != nil { + return err + } + // Baseline counters remain zero for checkpoint v1 compatibility; exact guard + // ownership is the authoritative append-only proof. + _, err = db.ExecContext(ctx, "INSERT INTO "+cp+" (id,marker,source_schema_sha256,source_fingerprint_sha256,phase,last_completed_end_id,seed_cutoff_id,generation,trigger_sql_mode,baseline_updates,baseline_deletes,updated_at) VALUES (1,?,?,?,'seed',0,?,0,?,0,0,CURRENT_TIMESTAMP(6))", marker, sourceHash, sourceFingerprint, maxID.Int64, sqlMode) + return err +} + +type columnSpec struct { + name, columnType, nullable, charset, collation, extra string +} + +var allowedLogCollations = []string{"utf8mb4_0900_ai_ci", "utf8mb4_unicode_ci"} + +var auditedLogColumns = []columnSpec{ + {"id", "bigint", "NO", "", "", "auto_increment"}, + {"user_id", "bigint", "YES", "", "", ""}, {"created_at", "bigint", "YES", "", "", ""}, {"type", "bigint", "YES", "", "", ""}, + {"content", "longtext", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, {"username", "varchar(191)", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, + {"token_name", "varchar(191)", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, {"model_name", "varchar(191)", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, + {"quota", "bigint", "YES", "", "", ""}, {"prompt_tokens", "bigint", "YES", "", "", ""}, {"completion_tokens", "bigint", "YES", "", "", ""}, + {"use_time", "bigint", "YES", "", "", ""}, {"is_stream", "tinyint(1)", "YES", "", "", ""}, {"channel_id", "bigint", "YES", "", "", ""}, + {"channel_name", "longtext", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, {"token_id", "bigint", "YES", "", "", ""}, + {"group", "varchar(191)", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, {"ip", "varchar(191)", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, + {"request_id", "varchar(64)", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, {"other", "longtext", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, + {"upstream_request_id", "varchar(128)", "YES", "utf8mb4", "utf8mb4_unicode_ci", ""}, +} + +func auditedLogsSchemaHash(ctx context.Context, q queryRowerQuerier, schema, table string) (string, error) { + hash, productionErr := auditedColumnsHash(ctx, q, schema, table, auditedLogColumns) + if productionErr == nil { + return hash, nil + } + hash, stagingErr := auditedColumnsHash(ctx, q, schema, table, auditedStagingLogColumns()) + if stagingErr == nil { + return hash, nil + } + return "", fmt.Errorf("table %s matches neither production nor staging audited logs schema: production=%v staging=%v", table, productionErr, stagingErr) +} + +func auditedStagingLogColumns() []columnSpec { + expected := slices.Clone(auditedLogColumns) + expected[19], expected[20] = expected[20], expected[19] + return expected +} + +func auditedColumnsHash(ctx context.Context, q queryRowerQuerier, schema, table string, expected []columnSpec) (string, error) { + rows, err := q.QueryContext(ctx, "SELECT COLUMN_NAME,COLUMN_TYPE,IS_NULLABLE,COALESCE(CHARACTER_SET_NAME,''),COALESCE(COLLATION_NAME,''),EXTRA FROM information_schema.columns WHERE table_schema=? AND table_name=? ORDER BY ORDINAL_POSITION", schema, table) + if err != nil { + return "", err + } + defer rows.Close() + h := sha256.New() + index := 0 + for rows.Next() { + var got columnSpec + if err := rows.Scan(&got.name, &got.columnType, &got.nullable, &got.charset, &got.collation, &got.extra); err != nil { + return "", err + } + if index >= len(expected) || !matchesAuditedColumn(got, expected[index]) { + return "", fmt.Errorf("table %s column %d differs from audited logs schema: got=%+v", table, index+1, got) + } + fmt.Fprintf(h, "%s|%s|%s|%s|%s|%s\n", got.name, got.columnType, got.nullable, got.charset, got.collation, got.extra) + index++ + } + if err := rows.Err(); err != nil { + return "", err + } + if index != len(expected) { + return "", fmt.Errorf("table %s has %d columns, audited schema requires %d", table, index, len(expected)) + } + return fmt.Sprintf("%x", h.Sum(nil)), nil +} + +func matchesAuditedColumn(got, want columnSpec) bool { + if got.charset == "utf8mb4" && want.charset == "utf8mb4" && slices.Contains(allowedLogCollations, got.collation) { + got.collation = want.collation + } + return got == want +} + +type queryRowerQuerier interface { + queryRower + QueryContext(context.Context, string, ...any) (*sql.Rows, error) +} + +func assertCheckpointSchema(ctx context.Context, db *sql.DB, c config) error { + expected := []string{"id|tinyint|NO", "marker|varchar(191)|NO", "source_schema_sha256|char(64)|NO", "source_fingerprint_sha256|char(64)|NO", "phase|varchar(32)|NO", "last_completed_end_id|bigint|NO", "seed_cutoff_id|bigint|NO", "final_cutoff_id|bigint|YES", "rollback_base_id|bigint|YES", "generation|bigint unsigned|NO", "trigger_sql_mode|text|NO", "ddl_operation|varchar(64)|YES", "ddl_nonce|varchar(64)|YES", "baseline_updates|bigint unsigned|NO", "baseline_deletes|bigint unsigned|NO", "updated_at|datetime(6)|NO"} + rows, err := db.QueryContext(ctx, "SELECT CONCAT(COLUMN_NAME,'|',COLUMN_TYPE,'|',IS_NULLABLE) FROM information_schema.columns WHERE table_schema=? AND table_name=? ORDER BY ORDINAL_POSITION", c.schema, c.checkpoint) + if err != nil { + return err + } + defer rows.Close() + var got []string + for rows.Next() { + var value string + if err := rows.Scan(&value); err != nil { + return err + } + got = append(got, value) + } + if err := rows.Err(); err != nil { + return err + } + if !slices.Equal(got, expected) { + return fmt.Errorf("checkpoint schema is not the exact owned v1 shape: got=%v", got) + } + return nil +} + +func assertRuntimeSafe(ctx context.Context, db *sql.DB, c config, state checkpoint, requireAppendOnly bool) error { + checkCtx, cancel := context.WithTimeout(ctx, c.statementTimeout) + defer cancel() + var variable string + var threads int + if err := db.QueryRowContext(checkCtx, "SHOW GLOBAL STATUS LIKE 'Threads_running'").Scan(&variable, &threads); err != nil { + return fmt.Errorf("read Threads_running: %w", err) + } + if threads > c.maxThreadsRunning { + return fmt.Errorf("Threads_running=%d exceeds stop threshold=%d", threads, c.maxThreadsRunning) + } + if requireAppendOnly { + if err := assertAppendOnlyGuards(checkCtx, db, c, state.triggerSQLMode); err != nil { + return err + } + } + return assertFrozenTableFingerprints(checkCtx, db, c, state) +} + +func assertAppendOnlyGuards(ctx context.Context, q queryRower, c config, sqlMode string) error { + for _, event := range []string{"update", "delete"} { + spec, err := expectedTriggerSpec(c, "guard_"+event, c.source, sqlMode) + if err != nil { + return err + } + got, exists, err := readTriggerSpec(ctx, q, c.schema, spec.name) + if err != nil || !exists || !triggerMatches(got, spec) { + return fmt.Errorf("append-only guard %s is missing or changed: exists=%t err=%v", spec.name, exists, err) + } + } + return nil +} + +func tableIdentity(ctx context.Context, db *sql.DB, schema, table string) (bool, string, error) { + var comment string + err := db.QueryRowContext(ctx, "SELECT TABLE_COMMENT FROM information_schema.tables WHERE table_schema=? AND table_name=?", schema, table).Scan(&comment) + if err == sql.ErrNoRows { + return false, "", nil + } + return err == nil, comment, err +} + +func schemaFingerprint(ctx context.Context, db *sql.DB, schema, table string) (string, error) { + h := sha256.New() + // Exclude mutable metadata such as AUTO_INCREMENT and TABLE_COMMENT, but bind + // every physical table property that CREATE TABLE LIKE must preserve. + var engine, rowFormat, collation, createOptions string + if err := db.QueryRowContext(ctx, "SELECT COALESCE(ENGINE,''),COALESCE(ROW_FORMAT,''),COALESCE(TABLE_COLLATION,''),COALESCE(CREATE_OPTIONS,'') FROM information_schema.tables WHERE table_schema=? AND table_name=?", schema, table).Scan(&engine, &rowFormat, &collation, &createOptions); err != nil { + return "", err + } + fmt.Fprintf(h, "table|%s|%s|%s|%s\n", engine, rowFormat, collation, createOptions) + rows, err := db.QueryContext(ctx, "SELECT COLUMN_NAME,ORDINAL_POSITION,COALESCE(COLUMN_DEFAULT,''),IS_NULLABLE,DATA_TYPE,COLUMN_TYPE,COALESCE(CHARACTER_SET_NAME,''),COALESCE(COLLATION_NAME,''),EXTRA,COALESCE(GENERATION_EXPRESSION,'') FROM information_schema.columns WHERE table_schema=? AND table_name=? ORDER BY ORDINAL_POSITION", schema, table) + if err != nil { + return "", err + } + for rows.Next() { + var a, b, c, d, e, f, g, i, j, k string + if err := rows.Scan(&a, &b, &c, &d, &e, &f, &g, &i, &j, &k); err != nil { + rows.Close() + return "", err + } + fmt.Fprintf(h, "%s|%s|%s|%s|%s|%s|%s|%s|%s|%s\n", a, b, c, d, e, f, g, i, j, k) + } + if err := rows.Close(); err != nil { + return "", err + } + rows, err = db.QueryContext(ctx, "SELECT INDEX_NAME,NON_UNIQUE,SEQ_IN_INDEX,COALESCE(COLUMN_NAME,''),COALESCE(COLLATION,''),COALESCE(SUB_PART,-1),NULLABLE,INDEX_TYPE,COALESCE(INDEX_COMMENT,''),IS_VISIBLE,COALESCE(EXPRESSION,'') FROM information_schema.statistics WHERE table_schema=? AND table_name=? ORDER BY INDEX_NAME,SEQ_IN_INDEX", schema, table) + if err != nil { + return "", err + } + for rows.Next() { + var a, b, c, d, e, f, g, i, j, k, l string + if err := rows.Scan(&a, &b, &c, &d, &e, &f, &g, &i, &j, &k, &l); err != nil { + rows.Close() + return "", err + } + fmt.Fprintf(h, "%s|%s|%s|%s|%s|%s|%s|%s|%s|%s|%s\n", a, b, c, d, e, f, g, i, j, k, l) + } + if err := rows.Close(); err != nil { + return "", err + } + rows, err = db.QueryContext(ctx, "SELECT COALESCE(PARTITION_NAME,''),COALESCE(SUBPARTITION_NAME,''),COALESCE(PARTITION_ORDINAL_POSITION,0),COALESCE(SUBPARTITION_ORDINAL_POSITION,0),COALESCE(PARTITION_METHOD,''),COALESCE(SUBPARTITION_METHOD,''),COALESCE(PARTITION_EXPRESSION,''),COALESCE(SUBPARTITION_EXPRESSION,''),COALESCE(PARTITION_DESCRIPTION,'') FROM information_schema.partitions WHERE table_schema=? AND table_name=? ORDER BY PARTITION_ORDINAL_POSITION,SUBPARTITION_ORDINAL_POSITION", schema, table) + if err != nil { + return "", err + } + for rows.Next() { + var a, b, c, d, e, f, g, i, j string + if err := rows.Scan(&a, &b, &c, &d, &e, &f, &g, &i, &j); err != nil { + rows.Close() + return "", err + } + fmt.Fprintf(h, "partition|%s|%s|%s|%s|%s|%s|%s|%s|%s\n", a, b, c, d, e, f, g, i, j) + } + if err := rows.Close(); err != nil { + return "", err + } + return fmt.Sprintf("%x", h.Sum(nil)), nil +} + +func assertFrozenTableFingerprints(ctx context.Context, db *sql.DB, c config, state checkpoint) error { + for _, table := range []string{c.source, c.target, c.old} { + exists, _, err := tableIdentity(ctx, db, c.schema, table) + if err != nil { + return err + } + if !exists { + continue + } + auditedHash, err := auditedLogsSchemaHash(ctx, db, c.schema, table) + if err != nil || auditedHash != state.sourceSchemaHash { + return fmt.Errorf("table %s audited schema changed got=%s want=%s err=%v", table, auditedHash, state.sourceSchemaHash, err) + } + fingerprint, err := schemaFingerprint(ctx, db, c.schema, table) + if err != nil || fingerprint != state.sourceFingerprint { + return fmt.Errorf("table %s full schema/index fingerprint changed got=%s want=%s err=%v", table, fingerprint, state.sourceFingerprint, err) + } + } + return nil +} + +func tableStateObserver(schema, table string, want bool) ddlObserver { + return func(ctx context.Context, conn *sql.Conn, _ string) (ddlState, error) { + var count int + if err := conn.QueryRowContext(ctx, "SELECT COUNT(*) FROM information_schema.tables WHERE table_schema=? AND table_name=?", schema, table).Scan(&count); err != nil { + return ddlUnknown, err + } + if (count == 1) == want { + return ddlPost, nil + } + return ddlPre, nil + } +} + +func tableCommentObserver(schema, table, marker string) ddlObserver { + return func(ctx context.Context, conn *sql.Conn, _ string) (ddlState, error) { + var comment string + err := conn.QueryRowContext(ctx, "SELECT TABLE_COMMENT FROM information_schema.tables WHERE table_schema=? AND table_name=?", schema, table).Scan(&comment) + if err == sql.ErrNoRows { + return ddlPre, nil + } + if err != nil { + return ddlUnknown, err + } + if comment == marker { + return ddlPost, nil + } + return ddlUnknown, nil + } +} + +func assertOwned(ctx context.Context, db *sql.DB, c config, table string) error { + var comment string + err := db.QueryRowContext(ctx, "SELECT TABLE_COMMENT FROM information_schema.tables WHERE table_schema=? AND table_name=?", c.schema, table).Scan(&comment) + if err != nil { + return err + } + if comment != ownershipMarker(c) { + return fmt.Errorf("ownership marker mismatch for %s", table) + } + return nil +} + +type checkpoint struct { + phase string + last, seed int64 + final, rollback sql.NullInt64 + generation uint64 + marker string + sourceSchemaHash string + sourceFingerprint string + triggerSQLMode string + ddlOperation sql.NullString + ddlNonce sql.NullString + baselineUpdates uint64 + baselineDeletes uint64 +} + +func loadCheckpoint(ctx context.Context, q interface { + QueryRowContext(context.Context, string, ...any) *sql.Row +}, c config) (checkpoint, error) { + cp, _ := qualified(c.schema, c.checkpoint) + var state checkpoint + err := q.QueryRowContext(ctx, "SELECT marker,source_schema_sha256,source_fingerprint_sha256,phase,last_completed_end_id,seed_cutoff_id,final_cutoff_id,rollback_base_id,generation,trigger_sql_mode,ddl_operation,ddl_nonce,baseline_updates,baseline_deletes FROM "+cp+" WHERE id=1").Scan(&state.marker, &state.sourceSchemaHash, &state.sourceFingerprint, &state.phase, &state.last, &state.seed, &state.final, &state.rollback, &state.generation, &state.triggerSQLMode, &state.ddlOperation, &state.ddlNonce, &state.baselineUpdates, &state.baselineDeletes) + if err == nil && state.marker != ownershipMarker(c) { + err = fmt.Errorf("checkpoint ownership marker mismatch") + } + return state, err +} + +func backfill(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + if err := assertOwned(ctx, db, c, c.target); err != nil { + return err + } + work, _, err := openVerifiedConn(ctx, db, c, "work") + if err != nil { + return err + } + defer work.Close() + if _, err = work.ExecContext(ctx, "SET SESSION TRANSACTION ISOLATION LEVEL READ COMMITTED"); err != nil { + return err + } + copySQL, _ := buildCopySQL(c) + source, _ := qualified(c.schema, c.source) + target, _ := qualified(c.schema, c.target) + cpSQL, _ := checkpointCASSQL(c) + initial, err := loadCheckpoint(ctx, work, c) + if err != nil { + return err + } + if err := assertRuntimeSafe(ctx, db, c, initial, true); err != nil { + return err + } + if initial.phase != c.phase { + if c.command != "reconcile" { + return fmt.Errorf("phase transition requires reconcile command: current=%s requested=%s", initial.phase, c.phase) + } + if err := transitionCheckpoint(ctx, work, c, initial, ev); err != nil { + return err + } + } + for { + if _, err := verifyPhysicalConn(ctx, work, c, "work-batch"); err != nil { + return err + } + state, err := loadCheckpoint(ctx, work, c) + if err != nil { + return err + } + upper := c.upperBound + if upper == 0 { + if state.phase == "fresh" && state.final.Valid { + upper = state.final.Int64 + } else { + upper = state.seed + } + } + var end sql.NullInt64 + err = work.QueryRowContext(ctx, "SELECT MAX(id) FROM (SELECT id FROM "+source+" FORCE INDEX(PRIMARY) WHERE id>? AND id<=? ORDER BY id LIMIT ?) x", state.last, upper, c.batchSize).Scan(&end) + if err != nil { + return err + } + if !end.Valid { + return nil + } + batchCtx, cancel := context.WithTimeout(ctx, c.statementTimeout) + tx, err := work.BeginTx(batchCtx, nil) + if err != nil { + cancel() + return err + } + if _, err = tx.ExecContext(batchCtx, copySQL, state.last, end.Int64, upper); err != nil { + _ = tx.Rollback() + cancel() + return err + } + if err = verifyWindow(batchCtx, tx, c, source, target, state.last, end.Int64); err != nil { + _ = tx.Rollback() + cancel() + return err + } + result, err := tx.ExecContext(batchCtx, cpSQL, c.phase, end.Int64, state.seed, state.final, state.generation) + if err != nil { + _ = tx.Rollback() + cancel() + return err + } + affected, _ := result.RowsAffected() + if affected != 1 { + _ = tx.Rollback() + cancel() + return fmt.Errorf("checkpoint CAS conflict generation=%d", state.generation) + } + if err = tx.Commit(); err != nil { + cancel() + return err + } + cancel() + if err := ev.emit("checkpoint_committed", map[string]any{"phase": c.phase, "end_id": end.Int64, "generation": state.generation + 1}); err != nil { + return err + } + select { + case <-ctx.Done(): + return context.Cause(ctx) + case <-time.After(c.batchDelay): + } + } +} + +func transitionCheckpoint(ctx context.Context, conn *sql.Conn, c config, state checkpoint, ev *evidence) error { + cp, _ := qualified(c.schema, c.checkpoint) + last, final, requireForward, err := checkpointTransitionPlan(state, c.phase, c.upperBound) + if err != nil { + return err + } + if requireForward { + spec, err := expectedTriggerSpec(c, "forward", c.source, state.triggerSQLMode) + if err != nil { + return err + } + got, exists, err := readTriggerSpec(ctx, conn, c.schema, spec.name) + if err != nil || !exists || !triggerMatches(got, spec) { + return fmt.Errorf("gap requires exact owned forward trigger: exists=%t err=%v", exists, err) + } + } + if err := ev.emit("checkpoint_transition_intent", map[string]any{"from": state.phase, "to": c.phase, "generation": state.generation, "reset_last_id": last}); err != nil { + return err + } + result, err := conn.ExecContext(ctx, "UPDATE "+cp+" SET phase=?,last_completed_end_id=?,final_cutoff_id=?,generation=generation+1,updated_at=CURRENT_TIMESTAMP(6) WHERE id=1 AND generation=?", c.phase, last, final, state.generation) + if err != nil { + return err + } + affected, _ := result.RowsAffected() + if affected != 1 { + return fmt.Errorf("checkpoint transition CAS conflict generation=%d", state.generation) + } + return ev.emit("checkpoint_transition_committed", map[string]any{"from": state.phase, "to": c.phase, "generation": state.generation + 1, "last_id": last}) +} + +func checkpointTransitionPlan(state checkpoint, next string, upper int64) (int64, sql.NullInt64, bool, error) { + final := state.final + switch { + case state.phase == "seed" && next == "gap": + if state.last < state.seed { + return 0, final, false, fmt.Errorf("seed phase is incomplete: last=%d seed=%d", state.last, state.seed) + } + if upper <= state.seed { + return 0, final, false, fmt.Errorf("gap upper-bound must exceed seed cutoff") + } + return state.seed, sql.NullInt64{Int64: upper, Valid: true}, true, nil + case state.phase == "gap" && next == "fresh": + if !state.final.Valid { + return 0, final, false, fmt.Errorf("fresh requires persisted final cutoff") + } + if state.last < state.final.Int64 { + return 0, final, false, fmt.Errorf("gap phase is incomplete: last=%d final=%d", state.last, state.final.Int64) + } + return 0, final, false, nil + case state.phase == "fresh" && next == "incremental": + if !state.final.Valid || upper <= state.final.Int64 { + return 0, final, false, fmt.Errorf("incremental upper-bound must exceed final cutoff") + } + if state.last < state.final.Int64 { + return 0, final, false, fmt.Errorf("fresh phase is incomplete: last=%d final=%d", state.last, state.final.Int64) + } + return state.final.Int64, final, false, nil + default: + return 0, final, false, fmt.Errorf("illegal checkpoint phase transition %s -> %s", state.phase, next) + } +} + +type queryRower interface { + QueryRowContext(context.Context, string, ...any) *sql.Row +} + +func verifyWindow(ctx context.Context, q queryRower, c config, source, target string, start, end int64) error { + predicate := retainedPredicate("s", c.channelIDs) + equal := rowEqualitySQL("s", "d") + queries := []string{ + "SELECT COUNT(*) FROM " + source + " s LEFT JOIN " + target + " d ON d.id=s.id WHERE s.id>? AND s.id<=? AND " + predicate + " AND d.id IS NULL", + "SELECT COUNT(*) FROM " + target + " d LEFT JOIN " + source + " s ON s.id=d.id WHERE d.id>? AND d.id<=? AND s.id IS NULL", + "SELECT COUNT(*) FROM " + source + " s JOIN " + target + " d ON d.id=s.id WHERE s.id>? AND s.id<=? AND " + predicate + " AND NOT (" + equal + ")", + "SELECT COUNT(*) FROM " + target + " d WHERE d.id>? AND d.id<=? AND NOT (" + retainedPredicate("d", c.channelIDs) + ")", + } + for i, query := range queries { + var count int + if err := q.QueryRowContext(ctx, query, start, end).Scan(&count); err != nil { + return err + } + if count != 0 { + return fmt.Errorf("window verification check %d failed count=%d range=(%d,%d]", i, count, start, end) + } + } + return nil +} + +func installForward(ctx context.Context, db *sql.DB, c config) error { + if c.triggerDefiner == "" { + return fmt.Errorf("trigger-definer is required") + } + if err := assertOwned(ctx, db, c, c.target); err != nil { + return err + } + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + if err := assertRuntimeSafe(ctx, db, c, state, true); err != nil { + return err + } + for _, event := range []string{"update", "delete"} { + query, _ := buildGuardTriggerSQL(c, event, c.source) + spec, err := expectedTriggerSpec(c, "guard_"+event, c.source, state.triggerSQLMode) + if err != nil { + return err + } + if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { + return err + } + } + query, err := buildForwardTriggerSQL(c) + if err != nil { + return err + } + spec, err := expectedTriggerSpec(c, "forward", c.source, state.triggerSQLMode) + if err != nil { + return err + } + return runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)) +} + +type triggerSpec struct { + schema, name, table, timing, event, action, definer, sqlMode string +} + +func expectedTriggerSpec(c config, kind, table, sqlMode string) (triggerSpec, error) { + name, err := triggerName(kind, c.batch) + if err != nil { + return triggerSpec{}, err + } + var statement, timing, event string + switch kind { + case "forward": + statement, err = buildForwardTriggerSQL(c) + timing, event = "AFTER", "INSERT" + case "reverse": + statement, err = buildStrictMirrorTriggerSQL(c, c.source, c.old) + timing, event = "AFTER", "INSERT" + case "guard_update": + statement, err = buildGuardTriggerSQL(c, "update", table) + timing, event = "BEFORE", "UPDATE" + case "guard_delete": + statement, err = buildGuardTriggerSQL(c, "delete", table) + timing, event = "BEFORE", "DELETE" + default: + return triggerSpec{}, fmt.Errorf("unsupported owned trigger kind %q", kind) + } + if err != nil { + return triggerSpec{}, err + } + marker := "FOR EACH ROW " + index := strings.Index(statement, marker) + if index < 0 { + return triggerSpec{}, fmt.Errorf("owned trigger SQL lacks action boundary") + } + return triggerSpec{schema: c.schema, name: name, table: table, timing: timing, event: event, action: statement[index+len(marker):], definer: c.triggerDefiner, sqlMode: sqlMode}, nil +} + +func exactTriggerObserver(spec triggerSpec, want bool) ddlObserver { + return func(ctx context.Context, conn *sql.Conn, ddlSQLMode string) (ddlState, error) { + got, exists, err := readTriggerSpec(ctx, conn, spec.schema, spec.name) + if err != nil { + return ddlUnknown, err + } + if !exists { + if want { + return ddlPre, nil + } + return ddlPost, nil + } + if !want { + return ddlPre, nil + } + if ddlSQLMode != spec.sqlMode || got.schema != spec.schema || got.name != spec.name || got.table != spec.table || got.timing != spec.timing || got.event != spec.event || got.definer != spec.definer || got.sqlMode != spec.sqlMode || !sameDDL(got.action, spec.action) { + return ddlUnknown, nil + } + return ddlPost, nil + } +} + +func readTriggerSpec(ctx context.Context, q queryRower, schema, name string) (triggerSpec, bool, error) { + var got triggerSpec + err := q.QueryRowContext(ctx, "SELECT TRIGGER_SCHEMA,TRIGGER_NAME,EVENT_OBJECT_TABLE,ACTION_TIMING,EVENT_MANIPULATION,ACTION_STATEMENT,DEFINER,SQL_MODE FROM information_schema.triggers WHERE trigger_schema=? AND trigger_name=?", schema, name).Scan(&got.schema, &got.name, &got.table, &got.timing, &got.event, &got.action, &got.definer, &got.sqlMode) + if err == sql.ErrNoRows { + return got, false, nil + } + return got, err == nil, err +} + +func triggerMatches(got, want triggerSpec) bool { + return got.schema == want.schema && got.name == want.name && got.table == want.table && got.timing == want.timing && got.event == want.event && got.definer == want.definer && got.sqlMode == want.sqlMode && sameDDL(got.action, want.action) +} + +func observeTopology(ctx context.Context, db *sql.DB, c config) (objectTopology, triggerTopology, error) { + var o objectTopology + for name, dst := range map[string]*bool{c.source: &o.source, c.target: &o.target, c.old: &o.old} { + var n int + if err := db.QueryRowContext(ctx, "SELECT COUNT(*) FROM information_schema.tables WHERE table_schema=? AND table_name=?", c.schema, name).Scan(&n); err != nil { + return o, triggerTopology{}, err + } + *dst = n == 1 + } + var t triggerTopology + checkpoint, err := loadCheckpoint(ctx, db, c) + if err != nil { + return o, t, err + } + status := classifyTopology(o) + for kind, dst := range map[string]*bool{"forward": &t.forward, "reverse": &t.reverse, "guard_update": &t.updateGuard, "guard_delete": &t.deleteGuard} { + name, _ := triggerName(kind, c.batch) + var n int + if err := db.QueryRowContext(ctx, "SELECT COUNT(*) FROM information_schema.triggers WHERE trigger_schema=? AND trigger_name=?", c.schema, name).Scan(&n); err != nil { + return o, t, err + } + *dst = n == 1 + if n == 1 { + got, exists, observeErr := readTriggerSpec(ctx, db, c.schema, name) + if observeErr != nil || !exists { + return o, t, fmt.Errorf("trigger %s exists but exact ownership is not proven: err=%v", name, observeErr) + } + table := c.source + switch { + case status == topologyPostCutover && kind == "forward": + table = c.old + case status == topologyPostCutover && (kind == "guard_update" || kind == "guard_delete"): + if got.table != c.source && got.table != c.old { + return o, t, fmt.Errorf("POST guard %s is attached to unexpected table %s", name, got.table) + } + table = got.table + case status == topologyPreCutover && (checkpoint.phase == "rollback-intent" || checkpoint.phase == "rollback-reconcile") && (kind == "guard_update" || kind == "guard_delete"): + if got.table != c.source && got.table != c.target { + return o, t, fmt.Errorf("rollback PRE guard %s is attached to unexpected table %s", name, got.table) + } + table = got.table + } + spec, specErr := expectedTriggerSpec(c, kind, table, checkpoint.triggerSQLMode) + if specErr != nil { + return o, t, specErr + } + if status == topologyPreCutover && (checkpoint.phase == "rollback-intent" || checkpoint.phase == "rollback-reconcile") && kind == "reverse" { + // RENAME moves the trigger with its source table but does not rewrite + // the stored action text. It is inactive on the compact target and is + // dropped before the new forward trigger is installed. + spec.table = c.target + } + if !triggerMatches(got, spec) { + return o, t, fmt.Errorf("trigger %s exists but exact ownership is not proven", name) + } + switch kind { + case "guard_update": + t.updateGuardTable = table + case "guard_delete": + t.deleteGuardTable = table + } + } + } + return o, t, nil +} + +func verify(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + wantTopology := classifyForVerify(state) + if wantTopology == topologyUnknown { + return fmt.Errorf("verify refuses transitional checkpoint phase %s; run recover", state.phase) + } + objects, _, err := observeTopology(ctx, db, c) + if err != nil { + return err + } + if got := classifyTopology(objects); got != wantTopology { + return fmt.Errorf("checkpoint expects %s topology, observed %s", wantTopology, got) + } + if err := assertRuntimeSafe(ctx, db, c, state, wantTopology == topologyPreCutover); err != nil { + return err + } + if wantTopology == topologyPostCutover { + // After the atomic RENAME, the owned shadow table becomes the live source; + // the original unowned source becomes the rollback table. + if err := assertOwned(ctx, db, c, ownedTableForTopology(c, wantTopology)); err != nil { + return err + } + if err := fullPostVerify(ctx, db, c); err != nil { + return err + } + return ev.emit("verify_passed", map[string]any{"topology": wantTopology}) + } + if err := assertOwned(ctx, db, c, ownedTableForTopology(c, wantTopology)); err != nil { + return err + } + source, _ := qualified(c.schema, c.source) + target, _ := qualified(c.schema, c.target) + upper := c.upperBound + if upper == 0 { + upper = state.last + } + start := int64(0) + for start < upper { + var end sql.NullInt64 + if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM (SELECT id FROM "+source+" WHERE id>? AND id<=? ORDER BY id LIMIT ?) x", start, upper, c.batchSize).Scan(&end); err != nil { + return err + } + if !end.Valid { + break + } + if err := verifyWindow(ctx, db, c, source, target, start, end.Int64); err != nil { + return err + } + start = end.Int64 + } + return ev.emit("verify_passed", map[string]any{"topology": wantTopology, "upper_bound": upper}) +} + +func classifyForVerify(state checkpoint) topologyStatus { + if state.phase == "rollback-gap" || state.phase == "rollback-ready" { + return topologyPostCutover + } + if state.phase == "ddl-intent" || state.phase == "rollback-intent" || state.phase == "rollback-reconcile" { + return topologyUnknown + } + return topologyPreCutover +} + +func ownedTableForTopology(c config, status topologyStatus) string { + if status == topologyPostCutover { + return c.source + } + return c.target +} + +func recover(ctx context.Context, db *sql.DB, c config, ev *evidence) error { + return recoverCutover(ctx, db, c, ev) +} + +func cleanup(ctx context.Context, db *sql.DB, c config) error { + o, _, err := observeTopology(ctx, db, c) + if err != nil { + return err + } + marker := ownershipMarker(c) + if c.confirmCleanup != marker { + return fmt.Errorf("confirm-cleanup must exactly equal ownership marker %q", marker) + } + if err := assertOwned(ctx, db, c, c.target); err != nil { + return err + } + plan, err := cleanupPlan(classifyTopology(o), true) + if err != nil { + return err + } + _ = plan + for _, kind := range []string{"forward", "guard_update", "guard_delete"} { + name, _ := triggerName(kind, c.batch) + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + spec, err := expectedTriggerSpec(c, kind, c.source, state.triggerSQLMode) + if err != nil { + return err + } + got, exists, err := readTriggerSpec(ctx, db, c.schema, name) + if err != nil { + return err + } + if !exists { + continue + } + if !triggerMatches(got, spec) { + return fmt.Errorf("refusing to drop non-owned trigger %s", name) + } + qname, _ := quoteIdentifier(name) + statement := "DROP TRIGGER " + qname + if err := runDDL(ctx, db, c, statement, exactTriggerObserver(spec, false)); err != nil { + return err + } + } + for _, table := range []string{c.target, c.checkpoint} { + if err := assertOwned(ctx, db, c, table); err != nil { + return err + } + qt, _ := qualified(c.schema, table) + if err := runDDL(ctx, db, c, "DROP TABLE "+qt, tableStateObserver(c.schema, table, false)); err != nil { + return err + } + } + return nil +} diff --git a/cmd/logs_slimming/sqlgen.go b/cmd/logs_slimming/sqlgen.go new file mode 100644 index 00000000000..c63d0ec0f2a --- /dev/null +++ b/cmd/logs_slimming/sqlgen.go @@ -0,0 +1,188 @@ +package main + +import ( + "crypto/sha256" + "fmt" + "math" + "strconv" + "strings" +) + +var logColumns = []string{ + "id", "user_id", "created_at", "type", "content", "username", "token_name", "model_name", + "quota", "prompt_tokens", "completion_tokens", "use_time", "is_stream", "channel_id", + "channel_name", "token_id", "`group`", "ip", "request_id", "other", "upstream_request_id", +} + +var textColumns = []string{"content", "username", "token_name", "model_name", "channel_name", "`group`", "ip", "request_id", "other", "upstream_request_id"} + +func auditedSchemaContractHash() string { + var b strings.Builder + for _, column := range auditedLogColumns { + fmt.Fprintf(&b, "%s|%s|%s|%s|%s|%s\n", column.name, column.columnType, column.nullable, column.charset, column.collation, column.extra) + } + return fmt.Sprintf("%x", sha256.Sum256([]byte(b.String()))) +} + +func qualified(schema, table string) (string, error) { + qs, err := quoteIdentifier(schema) + if err != nil { + return "", err + } + qt, err := quoteIdentifier(table) + if err != nil { + return "", err + } + return qs + "." + qt, nil +} + +func channelList(ids []int64) string { + parts := make([]string, len(ids)) + for i, id := range ids { + parts[i] = strconv.FormatInt(id, 10) + } + return strings.Join(parts, ",") +} + +func retainedPredicate(alias string, ids []int64) string { + return "NOT (" + filteredPredicate(alias, ids) + ")" +} + +func filteredPredicate(alias string, ids []int64) string { + prefix := "" + if alias != "" { + prefix = alias + "." + } + return fmt.Sprintf("COALESCE(%suser_id, -1) = 1 AND COALESCE(%stoken_id, 0) > 0 AND COALESCE(%schannel_id, -1) IN (%s)", prefix, prefix, prefix, channelList(ids)) +} + +func buildCopySQL(c config) (string, error) { + source, err := qualified(c.schema, c.source) + if err != nil { + return "", err + } + target, err := qualified(c.schema, c.target) + if err != nil { + return "", err + } + columns := strings.Join(logColumns, ", ") + return fmt.Sprintf("INSERT INTO %s (%s) SELECT %s FROM %s FORCE INDEX (PRIMARY) WHERE id > ? AND id <= ? AND id <= ? AND %s ON DUPLICATE KEY UPDATE id = %s.id", target, columns, columns, source, retainedPredicate("", c.channelIDs), target), nil +} + +func rowEqualitySQL(left, right string) string { + text := make(map[string]struct{}, len(textColumns)) + for _, column := range textColumns { + text[column] = struct{}{} + } + parts := make([]string, 0, len(logColumns)) + for _, column := range logColumns { + if _, ok := text[column]; ok { + parts = append(parts, fmt.Sprintf("BINARY %s.%s <=> BINARY %s.%s", left, column, right, column)) + } else { + parts = append(parts, fmt.Sprintf("%s.%s <=> %s.%s", left, column, right, column)) + } + } + return strings.Join(parts, " AND ") +} + +func triggerName(kind, batch string) (string, error) { + name := "logs_slim_" + kind + "_" + batch + if _, err := quoteIdentifier(name); err != nil { + return "", err + } + return name, nil +} + +func triggerDefiner(definer string) (string, error) { + parts := strings.Split(definer, "@") + if len(parts) != 2 || parts[0] == "" || parts[1] == "" || strings.ContainsAny(definer, "`'\";\\") { + return "", fmt.Errorf("trigger definer must be a simple user@host value") + } + return "`" + parts[0] + "`@`" + parts[1] + "`", nil +} + +func buildForwardTriggerSQL(c config) (string, error) { + name, err := triggerName("forward", c.batch) + if err != nil { + return "", err + } + qn, _ := quoteIdentifier(name) + source, err := qualified(c.schema, c.source) + if err != nil { + return "", err + } + target, err := qualified(c.schema, c.target) + if err != nil { + return "", err + } + definer, err := triggerDefiner(c.triggerDefiner) + if err != nil { + return "", err + } + values := make([]string, len(logColumns)) + for i, column := range logColumns { + values[i] = "NEW." + column + } + return fmt.Sprintf("CREATE DEFINER=%s TRIGGER %s AFTER INSERT ON %s FOR EACH ROW BEGIN IF %s THEN INSERT INTO %s (%s) VALUES (%s); END IF; END", definer, qn, source, retainedPredicate("NEW", c.channelIDs), target, strings.Join(logColumns, ", "), strings.Join(values, ", ")), nil +} + +func buildStrictMirrorTriggerSQL(c config, from, to string) (string, error) { + name, err := triggerName("reverse", c.batch) + if err != nil { + return "", err + } + qn, _ := quoteIdentifier(name) + qfrom, err := qualified(c.schema, from) + if err != nil { + return "", err + } + qto, err := qualified(c.schema, to) + if err != nil { + return "", err + } + definer, err := triggerDefiner(c.triggerDefiner) + if err != nil { + return "", err + } + values := make([]string, len(logColumns)) + for i, column := range logColumns { + values[i] = "NEW." + column + } + return fmt.Sprintf("CREATE DEFINER=%s TRIGGER %s AFTER INSERT ON %s FOR EACH ROW INSERT INTO %s (%s) VALUES (%s)", definer, qn, qfrom, qto, strings.Join(logColumns, ", "), strings.Join(values, ", ")), nil +} + +func buildGuardTriggerSQL(c config, event, table string) (string, error) { + kind := strings.ToLower(event) + if kind != "update" && kind != "delete" { + return "", fmt.Errorf("unsupported guard event %q", event) + } + name, err := triggerName("guard_"+kind, c.batch) + if err != nil { + return "", err + } + qn, _ := quoteIdentifier(name) + source, err := qualified(c.schema, table) + if err != nil { + return "", err + } + definer, err := triggerDefiner(c.triggerDefiner) + if err != nil { + return "", err + } + return fmt.Sprintf("CREATE DEFINER=%s TRIGGER %s BEFORE %s ON %s FOR EACH ROW SIGNAL SQLSTATE '45000' SET MYSQL_ERRNO=1644, MESSAGE_TEXT='logs slimming append-only guard'", definer, qn, strings.ToUpper(kind), source), nil +} + +func checkpointCASSQL(c config) (string, error) { + table, err := qualified(c.schema, c.checkpoint) + if err != nil { + return "", err + } + return "UPDATE " + table + " SET phase = ?, last_completed_end_id = ?, seed_cutoff_id = ?, final_cutoff_id = ?, generation = generation + 1, updated_at = CURRENT_TIMESTAMP(6) WHERE id = 1 AND generation = ?", nil +} + +func targetAutoIncrement(sourceNext, reserve uint64) (uint64, error) { + if reserve == 0 || sourceNext > math.MaxUint64-reserve { + return 0, fmt.Errorf("AUTO_INCREMENT reserve overflow") + } + return sourceNext + reserve, nil +} diff --git a/cmd/logs_slimming/topology.go b/cmd/logs_slimming/topology.go new file mode 100644 index 00000000000..7590869f8e6 --- /dev/null +++ b/cmd/logs_slimming/topology.go @@ -0,0 +1,91 @@ +package main + +import "fmt" + +type topologyStatus string + +const ( + topologyPreCutover topologyStatus = "PRE_CUTOVER" + topologyPostCutover topologyStatus = "POST_CUTOVER" + topologyUnknown topologyStatus = "UNKNOWN" +) + +type objectTopology struct { + source bool + target bool + old bool +} + +type triggerTopology struct { + forward bool + reverse bool + updateGuard bool + deleteGuard bool + updateGuardTable string + deleteGuardTable string +} + +func classifyTopology(t objectTopology) topologyStatus { + switch { + case t.source && t.target && !t.old: + return topologyPreCutover + case t.source && !t.target && t.old: + return topologyPostCutover + default: + return topologyUnknown + } +} + +type recoveryStep string + +const ( + stepRecordRollbackBase recoveryStep = "record-rollback-base" + stepDropForward recoveryStep = "drop-forward" + stepCreateReverse recoveryStep = "create-reverse" + stepReconcileRollbackGap recoveryStep = "reconcile-rollback-gap" +) + +func recoveryPlan(status topologyStatus, triggers triggerTopology) ([]recoveryStep, error) { + if triggers.forward && triggers.reverse { + return nil, fmt.Errorf("unsafe trigger cycle: forward and reverse triggers coexist") + } + switch status { + case topologyPreCutover: + if triggers.reverse { + return nil, fmt.Errorf("reverse trigger cannot exist before cutover") + } + return nil, nil + case topologyPostCutover: + var plan []recoveryStep + if triggers.forward { + plan = append(plan, stepRecordRollbackBase, stepDropForward, stepCreateReverse, stepReconcileRollbackGap) + return plan, nil + } + if !triggers.reverse { + plan = append(plan, stepRecordRollbackBase, stepCreateReverse, stepReconcileRollbackGap) + return plan, nil + } + return []recoveryStep{stepReconcileRollbackGap}, nil + default: + return nil, fmt.Errorf("cannot recover unknown object topology") + } +} + +type cleanupStep string + +const ( + cleanupDropForward cleanupStep = "drop-forward" + cleanupDropGuards cleanupStep = "drop-guards" + cleanupDropTarget cleanupStep = "drop-target" + cleanupDropCheckpoint cleanupStep = "drop-checkpoint" +) + +func cleanupPlan(status topologyStatus, ownershipConfirmed bool) ([]cleanupStep, error) { + if status != topologyPreCutover { + return nil, fmt.Errorf("cleanup only supports stable pre-cutover topology, got %s", status) + } + if !ownershipConfirmed { + return nil, fmt.Errorf("cleanup ownership is not confirmed") + } + return []cleanupStep{cleanupDropForward, cleanupDropGuards, cleanupDropTarget, cleanupDropCheckpoint}, nil +} From e9b9df9978e5c093558fe74a1ec3238d15528b43 Mon Sep 17 00:00:00 2001 From: slZhong <1542123803@qq.com> Date: Tue, 11 Aug 2026 17:30:31 +0800 Subject: [PATCH 2/3] Harden logs slimming cutover safety --- cmd/logs_slimming/core_test.go | 133 +++++++++++++++++------ cmd/logs_slimming/cutover.go | 187 ++++++++++++--------------------- cmd/logs_slimming/ddl.go | 68 ++++++++---- cmd/logs_slimming/evidence.go | 108 +++++++++++++++++-- cmd/logs_slimming/ops.go | 75 ++++++++++++- cmd/logs_slimming/sqlgen.go | 27 ++++- cmd/logs_slimming/topology.go | 21 ++-- 7 files changed, 425 insertions(+), 194 deletions(-) diff --git a/cmd/logs_slimming/core_test.go b/cmd/logs_slimming/core_test.go index e66d745bda8..008e7fe8f85 100644 --- a/cmd/logs_slimming/core_test.go +++ b/cmd/logs_slimming/core_test.go @@ -12,6 +12,10 @@ import ( "github.com/go-sql-driver/mysql" ) +type secretString string + +func (s secretString) String() string { return string(s) } + func validTestConfig() config { return config{ command: "backfill", @@ -263,6 +267,46 @@ func TestTriggerSQLUsesStrictInsertAndFrozenPredicate(t *testing.T) { } } +func TestUpdateGuardAllowsOnlyCompleteNoopDuplicateUpdate(t *testing.T) { + cfg := validTestConfig() + query, err := buildNamedGuardTriggerSQL(cfg, "future_guard_update", "update", cfg.target) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(query, "BEGIN IF NOT (") || !strings.Contains(query, "THEN SIGNAL SQLSTATE '45000'") { + t.Fatalf("update guard is not conditional: %s", query) + } + if got := strings.Count(query, "<=>"); got != len(logColumns) { + t.Fatalf("update guard compares %d fields, want %d: %s", got, len(logColumns), query) + } + for _, column := range logColumns { + if !strings.Contains(query, "OLD."+column+" <=>") && !strings.Contains(query, "BINARY OLD."+column+" <=>") { + t.Fatalf("update guard omits OLD/NEW equality for %s", column) + } + } + copySQL, err := buildCopySQL(cfg) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(copySQL, "ON DUPLICATE KEY UPDATE id = `newapi_staging`.`logs_compact_20260811`.id") { + t.Fatalf("backfill duplicate path is not a qualified no-op: %s", copySQL) + } + mirrorSQL, err := buildMirrorCopySQL(cfg, cfg.source, cfg.old) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(mirrorSQL, "ON DUPLICATE KEY UPDATE id=`newapi_staging`.`logs_old_20260811`.id") { + t.Fatalf("rollback mirror duplicate path is not a qualified no-op: %s", mirrorSQL) + } + deleteGuard, err := buildNamedGuardTriggerSQL(cfg, "future_guard_delete", "delete", cfg.target) + if err != nil { + t.Fatal(err) + } + if strings.Contains(deleteGuard, "IF NOT") || !strings.Contains(deleteGuard, "FOR EACH ROW SIGNAL") { + t.Fatalf("delete guard must remain unconditional: %s", deleteGuard) + } +} + func TestCheckpointCASIsGenerationGuarded(t *testing.T) { query, err := checkpointCASSQL(validTestConfig()) if err != nil { @@ -339,17 +383,20 @@ func TestPostCutoverRecoveryNeverAllowsTriggerCycle(t *testing.T) { func TestCleanupRefusesUnknownOrPostCutoverTopology(t *testing.T) { for _, status := range []topologyStatus{topologyUnknown, topologyPostCutover} { - if _, err := cleanupPlan(status, true); err == nil { + if _, err := cleanupPlan(status, triggerTopology{}, true); err == nil { t.Fatalf("cleanup accepted topology %s", status) } } - plan, err := cleanupPlan(topologyPreCutover, true) + plan, err := cleanupPlan(topologyPreCutover, triggerTopology{}, true) if err != nil { t.Fatal(err) } if len(plan) == 0 || plan[len(plan)-1] != cleanupDropCheckpoint { t.Fatalf("unexpected cleanup plan %v", plan) } + if _, err := cleanupPlan(topologyPreCutover, triggerTopology{reverse: true}, true); err == nil { + t.Fatal("cleanup accepted a reverse trigger") + } } func TestDDLKillerRequiresExactConnectionAndStatement(t *testing.T) { @@ -384,10 +431,21 @@ func TestCutoverAndRollbackRenameStatementsAreSymmetric(t *testing.T) { func TestEvidenceRedactsSecrets(t *testing.T) { sink := &memoryEvidence{} - e := newEvidence(sink, []string{"user:secret@tcp(host)/db", "secret"}) - e.emit("failed", map[string]any{"error": errors.New("dial user:secret@tcp(host)/db failed"), "dsn": "secret"}) + quotedSecret := `s"e\cret` + e := newEvidence(sink, []string{"user:secret@tcp(host)/db", "secret", quotedSecret}) + e.emit("failed", map[string]any{ + "error": errors.New("dial user:secret@tcp(host)/db failed"), + "nested": map[string]any{"slice": []any{"secret", []byte("secret")}}, + "struct": struct { + Token string `json:"token"` + }{Token: "secret"}, + "stringer": secretString("secret"), + "escaped": struct { + Token string `json:"token"` + }{Token: quotedSecret}, + }) output := sink.String() - if strings.Contains(output, "secret") || strings.Contains(output, "user:") { + if strings.Contains(output, "secret") || strings.Contains(output, "user:") || strings.Contains(output, quotedSecret) || strings.Contains(output, `s\"e\\cret`) { t.Fatalf("evidence leaked secret: %s", output) } if !strings.Contains(output, "") { @@ -404,6 +462,30 @@ func TestReserveRejectsOverflowAndUsesMinimum(t *testing.T) { } } +func TestStableDDLPostconditionUsesObservedStateAfterResolvedExecution(t *testing.T) { + if got := stableDDLPostcondition(true, []ddlState{ddlPost, ddlPost}); got != ddlPost { + t.Fatalf("stable POST classified as %s", got) + } + for _, states := range [][]ddlState{{ddlPost}, {ddlPost, ddlPre}, {ddlUnknown, ddlUnknown}} { + if got := stableDDLPostcondition(true, states); got != ddlUnknown { + t.Fatalf("unstable states %v classified as %s", states, got) + } + } + if got := stableDDLPostcondition(false, []ddlState{ddlPost, ddlPost}); got != ddlUnknown { + t.Fatalf("unresolved execution classified as %s", got) + } +} + +func TestDDLObservationContextNeverExtendsWatchdogDeadline(t *testing.T) { + deadline := time.Now().Add(100 * time.Millisecond) + ctx, cancel := boundedBackground(deadline, 250*time.Millisecond) + defer cancel() + got, ok := ctx.Deadline() + if !ok || got.After(deadline) { + t.Fatalf("observation deadline %v exceeds watchdog %v", got, deadline) + } +} + func TestCutoverUsesDoubleReserveBeforeFinalBarrier(t *testing.T) { got, err := plannedTargetAutoIncrement(1_000, 1_000_000) if err != nil || got != 2_001_000 { @@ -483,38 +565,27 @@ func TestPostTriggerPlanIsIdempotentAndNeverCycles(t *testing.T) { } } -func TestPostGuardPlanMovesOldGuardAndIsIdempotentOnLive(t *testing.T) { - plan, err := postGuardPlan(true, "logs_old_20260811", "logs", "logs_old_20260811") - if err != nil || len(plan) != 2 || plan[0] != postGuardDropOld || plan[1] != postGuardCreateLive { - t.Fatalf("old guard plan=%v err=%v", plan, err) - } - plan, err = postGuardPlan(false, "", "logs", "logs_old_20260811") - if err != nil || len(plan) != 1 || plan[0] != postGuardCreateLive { - t.Fatalf("missing guard plan=%v err=%v", plan, err) - } - plan, err = postGuardPlan(true, "logs", "logs", "logs_old_20260811") - if err != nil || len(plan) != 0 { - t.Fatalf("live guard plan=%v err=%v", plan, err) - } - if _, err := postGuardPlan(true, "foreign", "logs", "logs_old_20260811"); err == nil { - t.Fatal("foreign guard table accepted") - } -} - -func TestGuardSQLCanTargetPostCutoverLiveTable(t *testing.T) { +func TestFutureGuardsMoveAtomicallyWithCompactTable(t *testing.T) { cfg := validTestConfig() - query, err := buildGuardTriggerSQL(cfg, "update", cfg.source) + query, err := buildNamedGuardTriggerSQL(cfg, "future_guard_update", "update", cfg.target) if err != nil { t.Fatal(err) } - if !strings.Contains(query, "BEFORE UPDATE ON `newapi_staging`.`logs`") { - t.Fatalf("guard is not attached to active logs: %s", query) + if !strings.Contains(query, "logs_slim_future_guard_update_") || !strings.Contains(query, "BEFORE UPDATE ON `newapi_staging`.`logs_compact_20260811`") { + t.Fatalf("future guard is not preinstalled on compact target: %s", query) + } + pre := triggerTopology{ + updateGuard: true, deleteGuard: true, updateGuardTable: cfg.source, deleteGuardTable: cfg.source, + futureUpdateGuard: true, futureDeleteGuard: true, futureUpdateGuardTable: cfg.target, futureDeleteGuardTable: cfg.target, } - if guardsOnLive(triggerTopology{updateGuard: true, deleteGuard: true, updateGuardTable: cfg.old, deleteGuardTable: cfg.source}, cfg.source) { - t.Fatal("rollback-ready accepted a guard still attached to old") + if !preGuardsReady(pre, cfg) { + t.Fatal("complete pre-cutover guard topology rejected") } - if !guardsOnLive(triggerTopology{updateGuard: true, deleteGuard: true, updateGuardTable: cfg.source, deleteGuardTable: cfg.source}, cfg.source) { - t.Fatal("live guard topology rejected") + post := pre + post.updateGuardTable, post.deleteGuardTable = cfg.old, cfg.old + post.futureUpdateGuardTable, post.futureDeleteGuardTable = cfg.source, cfg.source + if !postGuardsReady(post, cfg) { + t.Fatal("future guards were not recognized on post-cutover live table") } } diff --git a/cmd/logs_slimming/cutover.go b/cmd/logs_slimming/cutover.go index 9d7e715a176..9395ca91862 100644 --- a/cmd/logs_slimming/cutover.go +++ b/cmd/logs_slimming/cutover.go @@ -194,7 +194,7 @@ func assertCutoverPreconditions(ctx context.Context, db *sql.DB, c config) (chec if err != nil { return state, err } - if classifyTopology(o) != topologyPreCutover || !triggers.forward || !triggers.updateGuard || !triggers.deleteGuard || triggers.reverse { + if classifyTopology(o) != topologyPreCutover || !triggers.forward || triggers.reverse || !preGuardsReady(triggers, c) { return state, fmt.Errorf("cutover trigger/topology precondition failed objects=%+v triggers=%+v", o, triggers) } if err := assertRuntimeSafe(ctx, db, c, state, true); err != nil { @@ -280,10 +280,21 @@ func showCreateAutoIncrement(ctx context.Context, q showCreateQueryer, schema, t return 0, err } match := autoIncrementPattern.FindStringSubmatch(create) - if len(match) != 2 { - return 0, fmt.Errorf("SHOW CREATE TABLE %s lacks AUTO_INCREMENT", table) + if len(match) == 2 { + return strconv.ParseUint(match[1], 10, 64) } - return strconv.ParseUint(match[1], 10, 64) + var next sql.NullInt64 + if err := q.QueryRowContext(ctx, "SELECT AUTO_INCREMENT FROM information_schema.tables WHERE table_schema=? AND table_name=?", schema, table).Scan(&next); err != nil { + return 0, err + } + if next.Valid && next.Int64 > 0 { + return uint64(next.Int64), nil + } + var maxID uint64 + if err := q.QueryRowContext(ctx, "SELECT COALESCE(MAX(id),0) FROM "+qualifiedTable).Scan(&maxID); err != nil { + return 0, err + } + return safeNext(maxID) } func renameStatement(c config) string { @@ -327,6 +338,12 @@ func renameWithMDLBarrier(ctx context.Context, db *sql.DB, c config, ev *evidenc if err != nil { return topologyUnknown, true, err } + ddlStarted := false + defer func() { + if !ddlStarted { + _ = ddlConn.Close() + } + }() control, _, err := openVerifiedConn(opCtx, db, c, "cutover-control") if err != nil { _ = ddlConn.Close() @@ -338,10 +355,19 @@ func renameWithMDLBarrier(ctx context.Context, db *sql.DB, c config, ev *evidenc return topologyUnknown, true, err } tagged := fmt.Sprintf("/*logs_slim batch=%s operation=%s nonce=%s*/ %s", c.batch, operation, nonce, statement) - if err := ev.emit("rename_barrier_intent", map[string]any{"nonce": nonce, "connection_id": ddlFacts.connectionID, "statement_sha256": fmt.Sprintf("%x", sha256.Sum256([]byte(statement)))}); err != nil { + statementHash := fmt.Sprintf("%x", sha256.Sum256([]byte(statement))) + emitPostcondition := func(result string, status topologyStatus, resolved bool, execErr, observeErr error) error { + return ev.emit("rename_postcondition", map[string]any{ + "operation": operation, "nonce": nonce, "statement_sha256": statementHash, + "result": result, "topology": status, "exec_resolved": resolved, + "exec_error": execErr, "observe_error": observeErr, + }) + } + if err := ev.emit("rename_barrier_intent", map[string]any{"nonce": nonce, "connection_id": ddlFacts.connectionID, "statement_sha256": statementHash}); err != nil { return topologyUnknown, true, err } done := make(chan error, 1) + ddlStarted = true go func() { _, execErr := ddlConn.ExecContext(opCtx, tagged) done <- execErr @@ -357,15 +383,21 @@ func renameWithMDLBarrier(ctx context.Context, db *sql.DB, c config, ev *evidenc select { case <-done: resolved = true - case <-time.After(ddlSettleGrace): + case <-time.After(minDuration(ddlSettleGrace, positiveRemaining(deadline))): } if resolved { _ = ddlConn.Close() } else { go reapDDLConn(done, ddlConn) } - observeDeadline := maxTime(deadline, time.Now().Add(ddlObservationGrace)) - status, topologyErr := observeStableObjectTopology(db, c, observeDeadline) + status, topologyErr := observeStableObjectTopology(db, c, deadline) + result := "known" + if topologyErr != nil || status == topologyUnknown || !resolved { + result = "unknown" + } + if evidenceErr := emitPostcondition(result, status, resolved, cause, topologyErr); evidenceErr != nil { + return topologyUnknown, resolved, fmt.Errorf("%w; rename postcondition evidence: %v", cause, evidenceErr) + } return status, resolved, fmt.Errorf("%w; exact_kill=%v barrier_rollback=%v topology=%v", cause, killErr, rollbackErr, topologyErr) } pendingDeadline := deadline.Add(-600 * time.Millisecond) @@ -419,15 +451,20 @@ func renameWithMDLBarrier(ctx context.Context, db *sql.DB, c config, ev *evidenc execErr = lateErr } _ = ddlConn.Close() - case <-time.After(ddlSettleGrace): + case <-time.After(minDuration(ddlSettleGrace, positiveRemaining(deadline))): go reapDDLConn(done, ddlConn) } } - observeDeadline := maxTime(deadline, time.Now().Add(ddlObservationGrace)) - status, topologyErr := observeStableObjectTopology(db, c, observeDeadline) + status, topologyErr := observeStableObjectTopology(db, c, deadline) if topologyErr != nil { + if evidenceErr := emitPostcondition("unknown", topologyUnknown, resolved, execErr, topologyErr); evidenceErr != nil { + return topologyUnknown, resolved, fmt.Errorf("rename=%v topology=%v evidence=%w", execErr, topologyErr, evidenceErr) + } return topologyUnknown, resolved, fmt.Errorf("rename=%v topology=%w", execErr, topologyErr) } + if evidenceErr := emitPostcondition("known", status, resolved, execErr, nil); evidenceErr != nil { + return topologyUnknown, resolved, fmt.Errorf("write rename postcondition evidence: %w", evidenceErr) + } return status, resolved, execErr } @@ -740,44 +777,8 @@ func ensurePreTriggersAfterRollback(ctx context.Context, db *sql.DB, c config, s return err } } - for _, event := range []string{"update", "delete"} { - _, current, err := observeTopology(ctx, db, c) - if err != nil { - return err - } - exists, table := current.updateGuard, current.updateGuardTable - if event == "delete" { - exists, table = current.deleteGuard, current.deleteGuardTable - } - if exists && table == c.target { - kind := "guard_" + event - spec, err := expectedTriggerSpec(c, kind, c.target, state.triggerSQLMode) - if err != nil { - return err - } - name, _ := triggerName(kind, c.batch) - quoted, _ := quoteIdentifier(name) - if err := runDDL(ctx, db, c, "DROP TRIGGER "+quoted, exactTriggerObserver(spec, false)); err != nil { - return err - } - exists = false - } - if !exists { - query, err := buildGuardTriggerSQL(c, event, c.source) - if err != nil { - return err - } - spec, err := expectedTriggerSpec(c, "guard_"+event, c.source, state.triggerSQLMode) - if err != nil { - return err - } - if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { - return err - } - } - } _, final, err := observeTopology(ctx, db, c) - if err != nil || !final.forward || final.reverse || !guardsOnLive(final, c.source) { + if err != nil || !final.forward || final.reverse || !preGuardsReady(final, c) { return fmt.Errorf("rollback PRE trigger stabilization incomplete: triggers=%+v err=%v", final, err) } return nil @@ -873,7 +874,7 @@ func assertPreCutoverTriggerTopology(ctx context.Context, db *sql.DB, c config) if err != nil { return err } - if classifyTopology(o) != topologyPreCutover || !t.forward || !t.updateGuard || !t.deleteGuard || t.reverse { + if classifyTopology(o) != topologyPreCutover || !t.forward || t.reverse || !preGuardsReady(t, c) { return fmt.Errorf("unsafe PRE trigger topology objects=%+v triggers=%+v", o, t) } return assertOwned(ctx, db, c, c.target) @@ -984,83 +985,25 @@ func ensurePostTriggers(ctx context.Context, db *sql.DB, c config, state checkpo return err } } - for _, event := range []string{"update", "delete"} { - if err := ensurePostGuard(ctx, db, c, state, event); err != nil { - return err - } - } _, final, err := observeTopology(ctx, db, c) - if err != nil || final.forward || !final.reverse || !guardsOnLive(final, c.source) { + if err != nil || final.forward || !final.reverse || !postGuardsReady(final, c) { return fmt.Errorf("POST trigger stabilization incomplete: triggers=%+v err=%v", final, err) } return nil } -type postGuardAction string - -const ( - postGuardDropOld postGuardAction = "drop-old" - postGuardCreateLive postGuardAction = "create-live" -) - -func postGuardPlan(exists bool, table, live, old string) ([]postGuardAction, error) { - switch { - case !exists: - return []postGuardAction{postGuardCreateLive}, nil - case table == live: - return nil, nil - case table == old: - return []postGuardAction{postGuardDropOld, postGuardCreateLive}, nil - default: - return nil, fmt.Errorf("guard is attached to unexpected table %q", table) - } -} - -func ensurePostGuard(ctx context.Context, db *sql.DB, c config, state checkpoint, event string) error { - _, t, err := observeTopology(ctx, db, c) - if err != nil { - return err - } - exists, table := t.updateGuard, t.updateGuardTable - if event == "delete" { - exists, table = t.deleteGuard, t.deleteGuardTable - } - actions, err := postGuardPlan(exists, table, c.source, c.old) - if err != nil { - return err - } - kind := "guard_" + event - for _, action := range actions { - switch action { - case postGuardDropOld: - spec, err := expectedTriggerSpec(c, kind, c.old, state.triggerSQLMode) - if err != nil { - return err - } - name, _ := triggerName(kind, c.batch) - quoted, _ := quoteIdentifier(name) - if err := runDDL(ctx, db, c, "DROP TRIGGER "+quoted, exactTriggerObserver(spec, false)); err != nil { - return err - } - case postGuardCreateLive: - query, err := buildGuardTriggerSQL(c, event, c.source) - if err != nil { - return err - } - spec, err := expectedTriggerSpec(c, kind, c.source, state.triggerSQLMode) - if err != nil { - return err - } - if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { - return err - } - } - } - return nil +func preGuardsReady(t triggerTopology, c config) bool { + return t.updateGuard && t.deleteGuard && + t.updateGuardTable == c.source && t.deleteGuardTable == c.source && + t.futureUpdateGuard && t.futureDeleteGuard && + t.futureUpdateGuardTable == c.target && t.futureDeleteGuardTable == c.target } -func guardsOnLive(t triggerTopology, live string) bool { - return t.updateGuard && t.deleteGuard && t.updateGuardTable == live && t.deleteGuardTable == live +func postGuardsReady(t triggerTopology, c config) bool { + return t.updateGuard && t.deleteGuard && + t.updateGuardTable == c.old && t.deleteGuardTable == c.old && + t.futureUpdateGuard && t.futureDeleteGuard && + t.futureUpdateGuardTable == c.source && t.futureDeleteGuardTable == c.source } type postTriggerAction string @@ -1087,8 +1030,10 @@ func reconcileRollbackGap(ctx context.Context, db *sql.DB, c config, ev *evidenc source, _ := qualified(c.schema, c.source) old, _ := qualified(c.schema, c.old) cp, _ := qualified(c.schema, c.checkpoint) - columns := strings.Join(logColumns, ", ") - copySQL := "INSERT INTO " + old + " (" + columns + ") SELECT " + columns + " FROM " + source + " WHERE id>? AND id<=? ON DUPLICATE KEY UPDATE id=" + old + ".id" + copySQL, err := buildMirrorCopySQL(c, c.source, c.old) + if err != nil { + return err + } var upper sql.NullInt64 if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+source).Scan(&upper); err != nil { return err @@ -1246,7 +1191,7 @@ func assertRollbackReady(ctx context.Context, db *sql.DB, c config, state checkp if err != nil { return err } - if classifyTopology(o) != topologyPostCutover || t.forward || !t.reverse || !guardsOnLive(t, c.source) { + if classifyTopology(o) != topologyPostCutover || t.forward || !t.reverse || !postGuardsReady(t, c) { return fmt.Errorf("ROLLBACK_READY topology mismatch objects=%+v triggers=%+v", o, t) } return nil diff --git a/cmd/logs_slimming/ddl.go b/cmd/logs_slimming/ddl.go index b7d59171170..569d80e9fe8 100644 --- a/cmd/logs_slimming/ddl.go +++ b/cmd/logs_slimming/ddl.go @@ -127,19 +127,43 @@ func runDDL(ctx context.Context, db *sql.DB, c config, statement string, observe go reapDDLConn(done, ddlConn) } - observeDeadline := maxTime(deadline, time.Now().Add(ddlObservationGrace)) + emitUnknown := func(states []ddlState, observeErr error) error { + fields := map[string]any{ + "operation": operation, "nonce": nonce, "statement_sha256": statementHash, + "state": fmt.Sprint(states), "result": "unknown", "exec_resolved": resolved, + "exec_error": execErr, "kill_error": killErr, "observe_error": observeErr, + } + if err := ev.emit("ddl_postcondition", fields); err != nil { + return fmt.Errorf("write unknown DDL postcondition evidence: %w", err) + } + return nil + } + observeDeadline := deadline states := make([]ddlState, 0, 2) for i := 0; i < 2; i++ { + if positiveRemaining(observeDeadline) <= time.Nanosecond { + err := fmt.Errorf("DDL total watchdog expired before postcondition") + if emitErr := emitUnknown(states, err); emitErr != nil { + return emitErr + } + return err + } observeCtx, cancel := boundedBackground(observeDeadline, 250*time.Millisecond) observer, _, openErr := openVerifiedConn(observeCtx, db, c, "observer") if openErr != nil { cancel() + if emitErr := emitUnknown(states, openErr); emitErr != nil { + return emitErr + } return openErr } state, observeErr := observe(observeCtx, observer, ddlFacts.sqlMode) _ = observer.Close() cancel() if observeErr != nil { + if emitErr := emitUnknown(states, observeErr); emitErr != nil { + return emitErr + } return observeErr } states = append(states, state) @@ -147,42 +171,43 @@ func runDDL(ctx context.Context, db *sql.DB, c config, statement string, observe time.Sleep(minDuration(25*time.Millisecond, positiveRemaining(deadline))) } } - if !resolved || killErr != nil { - if err := ev.emit("ddl_postcondition", map[string]any{"operation": operation, "nonce": nonce, "statement_sha256": statementHash, "state": fmt.Sprint(states), "result": "unknown", "exec_resolved": resolved, "kill_error": killErr}); err != nil { - return fmt.Errorf("write unresolved DDL postcondition evidence: %w", err) + outcome := stableDDLPostcondition(resolved, states) + if !resolved { + if err := emitUnknown(states, nil); err != nil { + return err } return fmt.Errorf("DDL outcome UNKNOWN: exec_resolved=%t kill=%v states=%v", resolved, killErr, states) } - if states[0] != states[1] || states[0] == ddlUnknown { - if err := ev.emit("ddl_postcondition", map[string]any{"operation": operation, "nonce": nonce, "statement_sha256": statementHash, "state": fmt.Sprint(states), "result": "unknown"}); err != nil { - return fmt.Errorf("write unknown DDL postcondition evidence: %w", err) + if outcome == ddlUnknown { + if err := emitUnknown(states, nil); err != nil { + return err } return fmt.Errorf("DDL postcondition is unstable or unknown: %v", states) } - if states[0] == ddlPost { - if err := ev.emit("ddl_postcondition", map[string]any{"operation": operation, "nonce": nonce, "statement_sha256": statementHash, "state": string(ddlPost), "elapsed_ms": time.Since(started).Milliseconds()}); err != nil { + if outcome == ddlPost { + if err := ev.emit("ddl_postcondition", map[string]any{"operation": operation, "nonce": nonce, "statement_sha256": statementHash, "state": string(ddlPost), "elapsed_ms": time.Since(started).Milliseconds(), "exec_error": execErr, "kill_warning": killErr}); err != nil { return fmt.Errorf("write DDL postcondition evidence: %w", err) } return nil } + if err := ev.emit("ddl_postcondition", map[string]any{"operation": operation, "nonce": nonce, "statement_sha256": statementHash, "state": string(ddlPre), "elapsed_ms": time.Since(started).Milliseconds(), "exec_error": execErr, "kill_warning": killErr}); err != nil { + return fmt.Errorf("write DDL PRE postcondition evidence: %w", err) + } if execErr != nil { return execErr } return fmt.Errorf("DDL returned but postcondition remained PRE") } -const ( - ddlSettleGrace = 500 * time.Millisecond - ddlObservationGrace = 750 * time.Millisecond -) - -func maxTime(a, b time.Time) time.Time { - if a.After(b) { - return a +func stableDDLPostcondition(resolved bool, states []ddlState) ddlState { + if !resolved || len(states) != 2 || states[0] != states[1] || states[0] == ddlUnknown { + return ddlUnknown } - return b + return states[0] } +const ddlSettleGrace = 200 * time.Millisecond + func killExactDDL(ctx context.Context, control *sql.Conn, c config, lockCtx context.Context, ddlFacts sessionFacts, tagged string) error { if _, err := verifyPhysicalConn(ctx, control, c, "control-watchdog"); err != nil { return err @@ -213,11 +238,10 @@ func reapDDLConn(done <-chan error, conn *sql.Conn) { } func boundedBackground(deadline time.Time, maximum time.Duration) (context.Context, context.CancelFunc) { - remaining := positiveRemaining(deadline) - if remaining > maximum { - remaining = maximum + if time.Until(deadline) <= maximum { + return context.WithDeadline(context.Background(), deadline) } - return context.WithTimeout(context.Background(), remaining) + return context.WithTimeout(context.Background(), maximum) } func positiveRemaining(deadline time.Time) time.Duration { diff --git a/cmd/logs_slimming/evidence.go b/cmd/logs_slimming/evidence.go index 5015fddbf9e..8f6775327d3 100644 --- a/cmd/logs_slimming/evidence.go +++ b/cmd/logs_slimming/evidence.go @@ -5,6 +5,7 @@ import ( "encoding/json" "fmt" "io" + "reflect" "strings" "sync" "time" @@ -44,6 +45,10 @@ func (e *evidence) emit(event string, fields map[string]any) error { data := bytes.TrimSuffix(buffer.Bytes(), []byte{'\n'}) if err != nil { data = []byte(fmt.Sprintf(`{"timestamp":%q,"event":"evidence_marshal_failed","error":%q}`, time.Now().UTC().Format(time.RFC3339Nano), e.redactString(err.Error()))) + } else { + // Final serialized-output redaction is a fail-safe for structs and other + // JSON-marshalable values not covered by the typed recursive cases below. + data = []byte(e.redactString(string(data))) } e.mu.Lock() defer e.mu.Unlock() @@ -52,13 +57,104 @@ func (e *evidence) emit(event string, fields map[string]any) error { } func (e *evidence) redactValue(value any) any { - switch typed := value.(type) { - case error: - return e.redactString(typed.Error()) - case string: - return e.redactString(typed) + return e.redactReflect(reflect.ValueOf(value), make(map[reflectVisit]struct{})) +} + +type reflectVisit struct { + typ reflect.Type + ptr uintptr +} + +func (e *evidence) redactReflect(value reflect.Value, seen map[reflectVisit]struct{}) any { + if !value.IsValid() { + return nil + } + if value.CanInterface() { + switch typed := value.Interface().(type) { + case error: + return e.redactString(typed.Error()) + case fmt.Stringer: + return e.redactString(typed.String()) + } + } + switch value.Kind() { + case reflect.Interface: + if value.IsNil() { + return nil + } + return e.redactReflect(value.Elem(), seen) + case reflect.Pointer: + if value.IsNil() { + return nil + } + visit := reflectVisit{typ: value.Type(), ptr: value.Pointer()} + if _, exists := seen[visit]; exists { + return "" + } + seen[visit] = struct{}{} + defer delete(seen, visit) + return e.redactReflect(value.Elem(), seen) + case reflect.String: + return e.redactString(value.String()) + case reflect.Bool: + return value.Bool() + case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: + return value.Int() + case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr: + return value.Uint() + case reflect.Float32, reflect.Float64: + return value.Float() + case reflect.Slice, reflect.Array: + if value.Type().Elem().Kind() == reflect.Uint8 { + bytes := make([]byte, value.Len()) + for i := range bytes { + bytes[i] = byte(value.Index(i).Uint()) + } + return e.redactString(string(bytes)) + } + items := make([]any, value.Len()) + for i := range items { + items[i] = e.redactReflect(value.Index(i), seen) + } + return items + case reflect.Map: + if value.IsNil() { + return nil + } + items := make(map[string]any, value.Len()) + iterator := value.MapRange() + for iterator.Next() { + key := e.redactString(fmt.Sprint(iterator.Key().Interface())) + items[key] = e.redactReflect(iterator.Value(), seen) + } + return items + case reflect.Struct: + items := make(map[string]any) + typ := value.Type() + for i := 0; i < value.NumField(); i++ { + fieldType := typ.Field(i) + field := value.Field(i) + if !fieldType.IsExported() || !field.CanInterface() { + continue + } + name := fieldType.Name + if tag := fieldType.Tag.Get("json"); tag != "" { + parts := strings.Split(tag, ",") + if parts[0] == "-" { + continue + } + if parts[0] != "" { + name = parts[0] + } + } + items[name] = e.redactReflect(field, seen) + } + return items default: - return typed + if value.CanInterface() { + return value.Interface() + } + return fmt.Sprint(value) } } diff --git a/cmd/logs_slimming/ops.go b/cmd/logs_slimming/ops.go index 4b5bf70f460..a19abe2c139 100644 --- a/cmd/logs_slimming/ops.go +++ b/cmd/logs_slimming/ops.go @@ -76,6 +76,10 @@ func preflight(ctx context.Context, db *sql.DB, c config, ev *evidence) error { } currentChannelIDs = append(currentChannelIDs, id) } + if err := rows.Err(); err != nil { + rows.Close() + return fmt.Errorf("iterate Codex channel snapshot: %w", err) + } if err := rows.Close(); err != nil { return fmt.Errorf("close Codex channel snapshot: %w", err) } @@ -266,6 +270,21 @@ func prepare(ctx context.Context, db *sql.DB, c config) error { if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { return err } + // A separately named guard is preinstalled on the compact table. RENAME + // moves it with that table, so the new live logs table is protected + // atomically instead of waiting for post-cutover trigger DDL. + futureKind := "future_guard_" + event + futureQuery, err := buildNamedGuardTriggerSQL(c, futureKind, event, c.target) + if err != nil { + return err + } + futureSpec, err := expectedTriggerSpec(c, futureKind, c.target, sqlMode) + if err != nil { + return err + } + if err := runDDL(ctx, db, c, futureQuery, exactTriggerObserver(futureSpec, true)); err != nil { + return err + } } statement := "CREATE TABLE " + cp + " (id TINYINT NOT NULL PRIMARY KEY, marker VARCHAR(191) NOT NULL, source_schema_sha256 CHAR(64) NOT NULL, source_fingerprint_sha256 CHAR(64) NOT NULL, phase VARCHAR(32) NOT NULL, last_completed_end_id BIGINT NOT NULL, seed_cutoff_id BIGINT NOT NULL, final_cutoff_id BIGINT NULL, rollback_base_id BIGINT NULL, generation BIGINT UNSIGNED NOT NULL, trigger_sql_mode TEXT NOT NULL, ddl_operation VARCHAR(64) NULL, ddl_nonce VARCHAR(64) NULL, baseline_updates BIGINT UNSIGNED NOT NULL, baseline_deletes BIGINT UNSIGNED NOT NULL, updated_at DATETIME(6) NOT NULL) ENGINE=InnoDB COMMENT='" + marker + "'" cpExists, cpComment, err := tableIdentity(ctx, db, c.schema, c.checkpoint) @@ -465,6 +484,10 @@ func schemaFingerprint(ctx context.Context, db *sql.DB, schema, table string) (s } fmt.Fprintf(h, "%s|%s|%s|%s|%s|%s|%s|%s|%s|%s\n", a, b, c, d, e, f, g, i, j, k) } + if err := rows.Err(); err != nil { + rows.Close() + return "", err + } if err := rows.Close(); err != nil { return "", err } @@ -480,6 +503,10 @@ func schemaFingerprint(ctx context.Context, db *sql.DB, schema, table string) (s } fmt.Fprintf(h, "%s|%s|%s|%s|%s|%s|%s|%s|%s|%s|%s\n", a, b, c, d, e, f, g, i, j, k, l) } + if err := rows.Err(); err != nil { + rows.Close() + return "", err + } if err := rows.Close(); err != nil { return "", err } @@ -495,6 +522,10 @@ func schemaFingerprint(ctx context.Context, db *sql.DB, schema, table string) (s } fmt.Fprintf(h, "partition|%s|%s|%s|%s|%s|%s|%s|%s|%s\n", a, b, c, d, e, f, g, i, j) } + if err := rows.Err(); err != nil { + rows.Close() + return "", err + } if err := rows.Close(); err != nil { return "", err } @@ -800,6 +831,18 @@ func installForward(ctx context.Context, db *sql.DB, c config) error { if err := runDDL(ctx, db, c, query, exactTriggerObserver(spec, true)); err != nil { return err } + futureKind := "future_guard_" + event + futureQuery, err := buildNamedGuardTriggerSQL(c, futureKind, event, c.target) + if err != nil { + return err + } + futureSpec, err := expectedTriggerSpec(c, futureKind, c.target, state.triggerSQLMode) + if err != nil { + return err + } + if err := runDDL(ctx, db, c, futureQuery, exactTriggerObserver(futureSpec, true)); err != nil { + return err + } } query, err := buildForwardTriggerSQL(c) if err != nil { @@ -835,6 +878,12 @@ func expectedTriggerSpec(c config, kind, table, sqlMode string) (triggerSpec, er case "guard_delete": statement, err = buildGuardTriggerSQL(c, "delete", table) timing, event = "BEFORE", "DELETE" + case "future_guard_update": + statement, err = buildNamedGuardTriggerSQL(c, kind, "update", table) + timing, event = "BEFORE", "UPDATE" + case "future_guard_delete": + statement, err = buildNamedGuardTriggerSQL(c, kind, "delete", table) + timing, event = "BEFORE", "DELETE" default: return triggerSpec{}, fmt.Errorf("unsupported owned trigger kind %q", kind) } @@ -899,7 +948,11 @@ func observeTopology(ctx context.Context, db *sql.DB, c config) (objectTopology, return o, t, err } status := classifyTopology(o) - for kind, dst := range map[string]*bool{"forward": &t.forward, "reverse": &t.reverse, "guard_update": &t.updateGuard, "guard_delete": &t.deleteGuard} { + for kind, dst := range map[string]*bool{ + "forward": &t.forward, "reverse": &t.reverse, + "guard_update": &t.updateGuard, "guard_delete": &t.deleteGuard, + "future_guard_update": &t.futureUpdateGuard, "future_guard_delete": &t.futureDeleteGuard, + } { name, _ := triggerName(kind, c.batch) var n int if err := db.QueryRowContext(ctx, "SELECT COUNT(*) FROM information_schema.triggers WHERE trigger_schema=? AND trigger_name=?", c.schema, name).Scan(&n); err != nil { @@ -915,6 +968,10 @@ func observeTopology(ctx context.Context, db *sql.DB, c config) (objectTopology, switch { case status == topologyPostCutover && kind == "forward": table = c.old + case status == topologyPostCutover && (kind == "future_guard_update" || kind == "future_guard_delete"): + table = c.source + case status == topologyPreCutover && (kind == "future_guard_update" || kind == "future_guard_delete"): + table = c.target case status == topologyPostCutover && (kind == "guard_update" || kind == "guard_delete"): if got.table != c.source && got.table != c.old { return o, t, fmt.Errorf("POST guard %s is attached to unexpected table %s", name, got.table) @@ -944,6 +1001,10 @@ func observeTopology(ctx context.Context, db *sql.DB, c config) (objectTopology, t.updateGuardTable = table case "guard_delete": t.deleteGuardTable = table + case "future_guard_update": + t.futureUpdateGuardTable = table + case "future_guard_delete": + t.futureDeleteGuardTable = table } } } @@ -1028,7 +1089,7 @@ func recover(ctx context.Context, db *sql.DB, c config, ev *evidence) error { } func cleanup(ctx context.Context, db *sql.DB, c config) error { - o, _, err := observeTopology(ctx, db, c) + o, triggers, err := observeTopology(ctx, db, c) if err != nil { return err } @@ -1039,18 +1100,22 @@ func cleanup(ctx context.Context, db *sql.DB, c config) error { if err := assertOwned(ctx, db, c, c.target); err != nil { return err } - plan, err := cleanupPlan(classifyTopology(o), true) + plan, err := cleanupPlan(classifyTopology(o), triggers, true) if err != nil { return err } _ = plan - for _, kind := range []string{"forward", "guard_update", "guard_delete"} { + for _, kind := range []string{"forward", "guard_update", "guard_delete", "future_guard_update", "future_guard_delete"} { name, _ := triggerName(kind, c.batch) state, err := loadCheckpoint(ctx, db, c) if err != nil { return err } - spec, err := expectedTriggerSpec(c, kind, c.source, state.triggerSQLMode) + table := c.source + if kind == "future_guard_update" || kind == "future_guard_delete" { + table = c.target + } + spec, err := expectedTriggerSpec(c, kind, table, state.triggerSQLMode) if err != nil { return err } diff --git a/cmd/logs_slimming/sqlgen.go b/cmd/logs_slimming/sqlgen.go index c63d0ec0f2a..f07021751e9 100644 --- a/cmd/logs_slimming/sqlgen.go +++ b/cmd/logs_slimming/sqlgen.go @@ -69,6 +69,19 @@ func buildCopySQL(c config) (string, error) { return fmt.Sprintf("INSERT INTO %s (%s) SELECT %s FROM %s FORCE INDEX (PRIMARY) WHERE id > ? AND id <= ? AND id <= ? AND %s ON DUPLICATE KEY UPDATE id = %s.id", target, columns, columns, source, retainedPredicate("", c.channelIDs), target), nil } +func buildMirrorCopySQL(c config, from, to string) (string, error) { + source, err := qualified(c.schema, from) + if err != nil { + return "", err + } + target, err := qualified(c.schema, to) + if err != nil { + return "", err + } + columns := strings.Join(logColumns, ", ") + return "INSERT INTO " + target + " (" + columns + ") SELECT " + columns + " FROM " + source + " WHERE id>? AND id<=? ON DUPLICATE KEY UPDATE id=" + target + ".id", nil +} + func rowEqualitySQL(left, right string) string { text := make(map[string]struct{}, len(textColumns)) for _, column := range textColumns { @@ -152,11 +165,15 @@ func buildStrictMirrorTriggerSQL(c config, from, to string) (string, error) { } func buildGuardTriggerSQL(c config, event, table string) (string, error) { + return buildNamedGuardTriggerSQL(c, "guard_"+strings.ToLower(event), event, table) +} + +func buildNamedGuardTriggerSQL(c config, triggerKind, event, table string) (string, error) { kind := strings.ToLower(event) if kind != "update" && kind != "delete" { return "", fmt.Errorf("unsupported guard event %q", event) } - name, err := triggerName("guard_"+kind, c.batch) + name, err := triggerName(triggerKind, c.batch) if err != nil { return "", err } @@ -169,7 +186,13 @@ func buildGuardTriggerSQL(c config, event, table string) (string, error) { if err != nil { return "", err } - return fmt.Sprintf("CREATE DEFINER=%s TRIGGER %s BEFORE %s ON %s FOR EACH ROW SIGNAL SQLSTATE '45000' SET MYSQL_ERRNO=1644, MESSAGE_TEXT='logs slimming append-only guard'", definer, qn, strings.ToUpper(kind), source), nil + signal := "SIGNAL SQLSTATE '45000' SET MYSQL_ERRNO=1644, MESSAGE_TEXT='logs slimming append-only guard'" + if kind == "update" { + // INSERT ... ON DUPLICATE KEY UPDATE id= executes BEFORE UPDATE. + // Permit only a complete OLD/NEW no-op; every real mutation is rejected. + return fmt.Sprintf("CREATE DEFINER=%s TRIGGER %s BEFORE UPDATE ON %s FOR EACH ROW BEGIN IF NOT (%s) THEN %s; END IF; END", definer, qn, source, rowEqualitySQL("OLD", "NEW"), signal), nil + } + return fmt.Sprintf("CREATE DEFINER=%s TRIGGER %s BEFORE DELETE ON %s FOR EACH ROW %s", definer, qn, source, signal), nil } func checkpointCASSQL(c config) (string, error) { diff --git a/cmd/logs_slimming/topology.go b/cmd/logs_slimming/topology.go index 7590869f8e6..6866f5475d8 100644 --- a/cmd/logs_slimming/topology.go +++ b/cmd/logs_slimming/topology.go @@ -17,12 +17,16 @@ type objectTopology struct { } type triggerTopology struct { - forward bool - reverse bool - updateGuard bool - deleteGuard bool - updateGuardTable string - deleteGuardTable string + forward bool + reverse bool + updateGuard bool + deleteGuard bool + updateGuardTable string + deleteGuardTable string + futureUpdateGuard bool + futureDeleteGuard bool + futureUpdateGuardTable string + futureDeleteGuardTable string } func classifyTopology(t objectTopology) topologyStatus { @@ -80,12 +84,15 @@ const ( cleanupDropCheckpoint cleanupStep = "drop-checkpoint" ) -func cleanupPlan(status topologyStatus, ownershipConfirmed bool) ([]cleanupStep, error) { +func cleanupPlan(status topologyStatus, triggers triggerTopology, ownershipConfirmed bool) ([]cleanupStep, error) { if status != topologyPreCutover { return nil, fmt.Errorf("cleanup only supports stable pre-cutover topology, got %s", status) } if !ownershipConfirmed { return nil, fmt.Errorf("cleanup ownership is not confirmed") } + if triggers.reverse { + return nil, fmt.Errorf("cleanup refuses a reverse trigger in pre-cutover topology") + } return []cleanupStep{cleanupDropForward, cleanupDropGuards, cleanupDropTarget, cleanupDropCheckpoint}, nil } From 35efe97643be06ec7679d05418178df566124021 Mon Sep 17 00:00:00 2001 From: slZhong <1542123803@qq.com> Date: Tue, 11 Aug 2026 18:25:36 +0800 Subject: [PATCH 3/3] Fix rollback filtered log cleanup recovery --- cmd/logs_slimming/core_test.go | 90 ++++++++++++++++++++++++++++++++++ cmd/logs_slimming/cutover.go | 74 +++++++++++++++++++++++++++- 2 files changed, 162 insertions(+), 2 deletions(-) diff --git a/cmd/logs_slimming/core_test.go b/cmd/logs_slimming/core_test.go index 008e7fe8f85..f00057ccb8e 100644 --- a/cmd/logs_slimming/core_test.go +++ b/cmd/logs_slimming/core_test.go @@ -589,6 +589,96 @@ func TestFutureGuardsMoveAtomicallyWithCompactTable(t *testing.T) { } } +func TestFilteredCleanupGuardPlanAllowsOnlyRollbackReconcile(t *testing.T) { + cfg := validTestConfig() + complete := triggerTopology{ + updateGuard: true, deleteGuard: true, updateGuardTable: cfg.source, deleteGuardTable: cfg.source, + futureUpdateGuard: true, futureDeleteGuard: true, futureUpdateGuardTable: cfg.target, futureDeleteGuardTable: cfg.target, + } + + drop, err := filteredCleanupGuardPlan("rollback-reconcile", complete, cfg) + if err != nil || !drop { + t.Fatalf("complete rollback cleanup topology rejected: drop=%t err=%v", drop, err) + } + + missingDelete := complete + missingDelete.futureDeleteGuard = false + missingDelete.futureDeleteGuardTable = "" + drop, err = filteredCleanupGuardPlan("rollback-reconcile", missingDelete, cfg) + if err != nil || drop { + t.Fatalf("recoverable missing future DELETE guard rejected: drop=%t err=%v", drop, err) + } + + for _, phase := range []string{"fresh", "rollback-intent"} { + if _, err := filteredCleanupGuardPlan(phase, missingDelete, cfg); err == nil { + t.Fatalf("phase %s accepted a missing future DELETE guard", phase) + } + } +} + +func TestFilteredCleanupGuardPlanRejectsUnsafeGuardTopology(t *testing.T) { + cfg := validTestConfig() + complete := triggerTopology{ + updateGuard: true, deleteGuard: true, updateGuardTable: cfg.source, deleteGuardTable: cfg.source, + futureUpdateGuard: true, futureDeleteGuard: true, futureUpdateGuardTable: cfg.target, futureDeleteGuardTable: cfg.target, + } + + tests := map[string]func(*triggerTopology){ + "future delete on source": func(topology *triggerTopology) { topology.futureDeleteGuardTable = cfg.source }, + "future delete elsewhere": func(topology *triggerTopology) { topology.futureDeleteGuardTable = "other_logs" }, + "missing future delete with stale table": func(topology *triggerTopology) { + topology.futureDeleteGuard = false + topology.futureDeleteGuardTable = cfg.target + }, + "missing source update": func(topology *triggerTopology) { topology.updateGuard = false }, + "missing source delete": func(topology *triggerTopology) { topology.deleteGuard = false }, + "missing future update": func(topology *triggerTopology) { topology.futureUpdateGuard = false }, + } + for name, mutate := range tests { + t.Run(name, func(t *testing.T) { + topology := complete + mutate(&topology) + if _, err := filteredCleanupGuardPlan("rollback-reconcile", topology, cfg); err == nil { + t.Fatalf("unsafe topology accepted: %+v", topology) + } + }) + } +} + +func TestFilteredCleanupGuardRestoreUsesPersistedTriggerIdentity(t *testing.T) { + cfg := validTestConfig() + const persistedSQLMode = "STRICT_TRANS_TABLES,NO_ENGINE_SUBSTITUTION" + spec, err := expectedTriggerSpec(cfg, "future_guard_delete", cfg.target, persistedSQLMode) + if err != nil { + t.Fatal(err) + } + name, err := triggerName("future_guard_delete", cfg.batch) + if err != nil { + t.Fatal(err) + } + if spec.name != name || spec.table != cfg.target || spec.sqlMode != persistedSQLMode || spec.definer != cfg.triggerDefiner { + t.Fatalf("restore spec does not preserve trigger identity: %+v", spec) + } + query, err := buildNamedGuardTriggerSQL(cfg, "future_guard_delete", "delete", cfg.target) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(query, "TRIGGER `"+name+"`") || !strings.Contains(query, "BEFORE DELETE ON `"+cfg.schema+"`.`"+cfg.target+"`") { + t.Fatalf("restore SQL targets the wrong trigger/table: %s", query) + } +} + +func TestFilteredCleanupGuardRestoreReturnsToStrictPreGuards(t *testing.T) { + cfg := validTestConfig() + topology := triggerTopology{ + updateGuard: true, deleteGuard: true, updateGuardTable: cfg.source, deleteGuardTable: cfg.source, + futureUpdateGuard: true, futureDeleteGuard: true, futureUpdateGuardTable: cfg.target, futureDeleteGuardTable: cfg.target, + } + if !preGuardsReady(topology, cfg) { + t.Fatal("restored cleanup guard topology is not strict pre-cutover topology") + } +} + func TestNextPostVerifyEndUsesSmallestValidBoundedEndpoint(t *testing.T) { tests := []struct { name string diff --git a/cmd/logs_slimming/cutover.go b/cmd/logs_slimming/cutover.go index 9395ca91862..e3e2e3b53ae 100644 --- a/cmd/logs_slimming/cutover.go +++ b/cmd/logs_slimming/cutover.go @@ -4,6 +4,7 @@ import ( "context" "crypto/sha256" "database/sql" + "errors" "fmt" "math" "regexp" @@ -686,7 +687,54 @@ func stabilizePreRollback(ctx context.Context, db *sql.DB, c config, ev *evidenc return assertPreCutoverTriggerTopology(ctx, db, c) } -func removeFilteredRollbackRows(ctx context.Context, db *sql.DB, c config, ev *evidence, verifiedFloor int64) error { +func removeFilteredRollbackRows(ctx context.Context, db *sql.DB, c config, ev *evidence, verifiedFloor int64) (retErr error) { + state, err := loadCheckpoint(ctx, db, c) + if err != nil { + return err + } + objects, topology, err := observeTopology(ctx, db, c) + if err != nil { + return err + } + if classifyTopology(objects) != topologyPreCutover || !topology.forward || topology.reverse { + return fmt.Errorf("rollback filtered cleanup requires stable PRE topology: objects=%+v triggers=%+v", objects, topology) + } + dropGuard, err := filteredCleanupGuardPlan(state.phase, topology, c) + if err != nil { + return err + } + guardSpec, err := expectedTriggerSpec(c, "future_guard_delete", c.target, state.triggerSQLMode) + if err != nil { + return err + } + if dropGuard { + name, _ := triggerName("future_guard_delete", c.batch) + quoted, _ := quoteIdentifier(name) + if err := runDDL(ctx, db, c, "DROP TRIGGER "+quoted, exactTriggerObserver(guardSpec, false)); err != nil { + return err + } + } + // The compact target is inactive throughout rollback-reconcile. Its DELETE + // guard may be absent only during this cleanup, and is always rebuilt before + // the checkpoint can leave rollback-reconcile. A crash after DROP is handled + // by the same idempotent path on the next recover. + defer func() { + query, err := buildNamedGuardTriggerSQL(c, "future_guard_delete", "delete", c.target) + if err == nil { + err = runDDL(ctx, db, c, query, exactTriggerObserver(guardSpec, true)) + } + if err == nil { + _, topology, observeErr := observeTopology(ctx, db, c) + if observeErr != nil { + err = observeErr + } else if !preGuardsReady(topology, c) { + err = fmt.Errorf("future DELETE guard restore did not recover exact pre-cutover guards: %+v", topology) + } + } + if err != nil { + retErr = errors.Join(retErr, fmt.Errorf("restore future DELETE guard: %w", err)) + } + }() target, _ := qualified(c.schema, c.target) var upper sql.NullInt64 if err := db.QueryRowContext(ctx, "SELECT MAX(id) FROM "+target).Scan(&upper); err != nil { @@ -744,6 +792,24 @@ func removeFilteredRollbackRows(ctx context.Context, db *sql.DB, c config, ev *e return nil } +func filteredCleanupGuardPlan(phase string, topology triggerTopology, c config) (bool, error) { + if phase != "rollback-reconcile" { + return false, fmt.Errorf("future DELETE guard may be suspended only during rollback-reconcile") + } + if !rollbackCleanupGuardsReady(topology, c) { + return false, fmt.Errorf("rollback filtered cleanup guard topology is unsafe: %+v", topology) + } + return topology.futureDeleteGuard, nil +} + +func rollbackCleanupGuardsReady(t triggerTopology, c config) bool { + return t.updateGuard && t.deleteGuard && + t.updateGuardTable == c.source && t.deleteGuardTable == c.source && + t.futureUpdateGuard && t.futureUpdateGuardTable == c.target && + (t.futureDeleteGuard && t.futureDeleteGuardTable == c.target || + !t.futureDeleteGuard && t.futureDeleteGuardTable == "") +} + func ensurePreTriggersAfterRollback(ctx context.Context, db *sql.DB, c config, state checkpoint) error { o, t, err := observeTopology(ctx, db, c) if err != nil { @@ -778,7 +844,11 @@ func ensurePreTriggersAfterRollback(ctx context.Context, db *sql.DB, c config, s } } _, final, err := observeTopology(ctx, db, c) - if err != nil || !final.forward || final.reverse || !preGuardsReady(final, c) { + guardsReady := preGuardsReady(final, c) + if state.phase == "rollback-reconcile" { + guardsReady = rollbackCleanupGuardsReady(final, c) + } + if err != nil || !final.forward || final.reverse || !guardsReady { return fmt.Errorf("rollback PRE trigger stabilization incomplete: triggers=%+v err=%v", final, err) } return nil