From 88f8ade6242b8aefb09d7072f0de62d2da084ee3 Mon Sep 17 00:00:00 2001 From: Ian Littman Date: Thu, 22 Jan 2026 15:00:21 -0600 Subject: [PATCH] Add step-based and intra-step framework for migration progress (#38556) Resolves #35916. For example: ```go return withSteps([]migrationStep{ basicMigrationStep("SELECT NOW()", "couldn't select from hosts"), incrementalMigrationStep(func(tx *sql.Tx) (uint64, error) { return 25, nil }, func(tx *sql.Tx, increment incrementCountFn) error { for range 25 { time.Sleep(time.Second) increment() } return nil }), }, tx) ``` gets you ``` 2026/01/20 17:16:30 [2026-01-20] Test Migration Step 1 of 2 Step 2 of 2 16% complete 36% complete 56% complete 76% complete 96% complete 100% complete Migrations completed. ``` No need to use this on short migrations, but we can throw this wherever on longer migrations, and progress display and upgdate frequency can be adjusted independent of the migrations themselves. # Checklist for submitter If some of the following don't apply, delete the relevant line. - [ ] Changes file added for user-visible changes in `changes/`, `orbit/changes/` or `ee/fleetd-chrome/changes`. See [Changes files](https://github.com/fleetdm/fleet/blob/main/docs/Contributing/guides/committing-changes.md#changes-files) for more information. - [x] Input data is properly validated, `SELECT *` is avoided, SQL injection is prevented (using placeholders for values in statements) ## Testing - [x] Added/updated automated tests - [x] QA'd all new/changed functionality manually --- .../mysql/migrations/tables/migration.go | 88 +++- .../tables/migration_helpers_test.go | 399 ++++++++++++++++++ 2 files changed, 483 insertions(+), 4 deletions(-) create mode 100644 server/datastore/mysql/migrations/tables/migration_helpers_test.go diff --git a/server/datastore/mysql/migrations/tables/migration.go b/server/datastore/mysql/migrations/tables/migration.go index 8908dd6cf7..ec30f56485 100644 --- a/server/datastore/mysql/migrations/tables/migration.go +++ b/server/datastore/mysql/migrations/tables/migration.go @@ -4,7 +4,11 @@ import ( "database/sql" "encoding/json" "fmt" + "io" + "os" "strings" + "sync/atomic" + "time" "github.com/fleetdm/fleet/v4/server/fleet" "github.com/fleetdm/fleet/v4/server/goose" @@ -14,14 +18,90 @@ import ( var MigrationClient = goose.New("migration_status_tables", goose.MySqlDialect{}) +// can override in tests +var outputTo io.Writer = os.Stderr +var progressInterval = time.Second * 5 + +type migrationStep func(tx *sql.Tx) error + +func basicMigrationStep(statement string, errorMessage string) migrationStep { + return func(tx *sql.Tx) error { + _, err := tx.Exec(statement) + return errors.Wrap(err, errorMessage) + } +} + +type getTotalCountFn func(tx *sql.Tx) (uint64, error) +type incrementCountFn func() +type executeWithProgressFn func(tx *sql.Tx, increment incrementCountFn) error + +func incrementalMigrationStep(count getTotalCountFn, execute executeWithProgressFn) migrationStep { + return func(tx *sql.Tx) error { + total, err := count(tx) + if err != nil { + return err + } + if total == 0 { // skip no-ops to avoid divide by zero + return nil + } + + atomicCurrent := atomic.Uint64{} + + // Every five seconds, echo the % progress of the executor + // Since we output once the migration step is complete, we need an extra channel to indicate when both the step + // and the "step complete" output are com0plete + stepComplete := make(chan struct{}) + outputComplete := make(chan struct{}) + go func() { + ticker := time.NewTicker(progressInterval) + defer ticker.Stop() + defer close(outputComplete) + for { + select { + case <-stepComplete: + _, _ = fmt.Fprint(outputTo, " 100% complete\n") + return + case <-ticker.C: + current := atomicCurrent.Load() + if current == total { + _, _ = fmt.Fprint(outputTo, " Almost done...\n") + } else { + _, _ = fmt.Fprintf(outputTo, " %d%% complete\n", (100*current)/total) + } + } + } + }() + + err = execute(tx, func() { + atomicCurrent.Add(1) + }) + close(stepComplete) + <-outputComplete // Wait for the goroutine to complete + return err + } +} + +func withSteps(steps []migrationStep, tx *sql.Tx) error { + stepCount := len(steps) + for i, step := range steps { + if stepCount > 1 { + _, _ = fmt.Fprintf(outputTo, " Step %d of %d\n", i+1, stepCount) + } + if err := step(tx); err != nil { + return err + } + } + return nil +} + func fkExists(tx *sql.Tx, table, name string) bool { var count int err := tx.QueryRow(` SELECT COUNT(1) FROM information_schema.REFERENTIAL_CONSTRAINTS -WHERE CONSTRAINT_SCHEMA = DATABASE() +WHERE CONSTRAINT_SCHEMA = DATABASE() AND TABLE_NAME = ? -AND CONSTRAINT_NAME = ? +AND CONSTRAINT_NAME = ? `, table, name).Scan(&count) if err != nil { return false @@ -35,9 +115,9 @@ func constraintExists(tx *sql.Tx, table, name string) bool { err := tx.QueryRow(` SELECT COUNT(1) FROM information_schema.TABLE_CONSTRAINTS -WHERE CONSTRAINT_SCHEMA = DATABASE() +WHERE CONSTRAINT_SCHEMA = DATABASE() AND TABLE_NAME = ? -AND CONSTRAINT_NAME = ? +AND CONSTRAINT_NAME = ? `, table, name).Scan(&count) if err != nil { return false diff --git a/server/datastore/mysql/migrations/tables/migration_helpers_test.go b/server/datastore/mysql/migrations/tables/migration_helpers_test.go new file mode 100644 index 0000000000..7f8fa0be92 --- /dev/null +++ b/server/datastore/mysql/migrations/tables/migration_helpers_test.go @@ -0,0 +1,399 @@ +package tables + +import ( + "bytes" + "database/sql" + "errors" + "io" + "strings" + "testing" + "time" + + "github.com/DATA-DOG/go-sqlmock" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestBasicMigrationStep(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + t.Run("success", func(t *testing.T) { + mock.ExpectBegin() + mock.ExpectExec("ALTER TABLE foo ADD COLUMN bar INT").WillReturnResult(sqlmock.NewResult(0, 0)) + + tx, err := db.Begin() + require.NoError(t, err) + + step := basicMigrationStep("ALTER TABLE foo ADD COLUMN bar INT", "failed to add column") + err = step(tx) + require.NoError(t, err) + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("error", func(t *testing.T) { + mock.ExpectBegin() + mock.ExpectExec("ALTER TABLE foo ADD COLUMN bar INT").WillReturnError(errors.New("syntax error")) + + tx, err := db.Begin() + require.NoError(t, err) + + step := basicMigrationStep("ALTER TABLE foo ADD COLUMN bar INT", "failed to add column") + err = step(tx) + require.Error(t, err) + assert.Contains(t, err.Error(), "failed to add column") + assert.Contains(t, err.Error(), "syntax error") + + require.NoError(t, mock.ExpectationsWereMet()) + }) +} + +func TestIncrementalMigrationStep(t *testing.T) { + // Save original values and restore after test + originalOutputTo := outputTo + originalProgressInterval := progressInterval + defer func() { + outputTo = originalOutputTo + progressInterval = originalProgressInterval + }() + + t.Run("zero count skips execution", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectBegin() + + tx, err := db.Begin() + require.NoError(t, err) + + executeCalled := false + step := incrementalMigrationStep( + func(tx *sql.Tx) (uint64, error) { + return 0, nil + }, + func(tx *sql.Tx, increment incrementCountFn) error { + executeCalled = true + return nil + }, + ) + + err = step(tx) + require.NoError(t, err) + assert.False(t, executeCalled, "executor should not be called when count is 0") + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("count error is returned", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectBegin() + + tx, err := db.Begin() + require.NoError(t, err) + + var wasExecutorCalled bool + expectedErr := errors.New("count query failed") + step := incrementalMigrationStep( + func(tx *sql.Tx) (uint64, error) { + return 0, expectedErr + }, + func(tx *sql.Tx, increment incrementCountFn) error { + wasExecutorCalled = true + return nil + }, + ) + + err = step(tx) + require.Error(t, err) + assert.Equal(t, expectedErr, err) + require.False(t, wasExecutorCalled, "executor should not be called when count call errors") + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("executor error is returned", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectBegin() + + tx, err := db.Begin() + require.NoError(t, err) + + expectedErr := errors.New("executor failed") + step := incrementalMigrationStep( + func(tx *sql.Tx) (uint64, error) { + return 5, nil + }, + func(tx *sql.Tx, increment incrementCountFn) error { + return expectedErr + }, + ) + + err = step(tx) + require.Error(t, err) + assert.Equal(t, expectedErr, err) + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("progress output", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectBegin() + + tx, err := db.Begin() + require.NoError(t, err) + + // Override output and progress interval for testing + var buf bytes.Buffer + outputTo = &buf + progressInterval = 10 * time.Millisecond + + step := incrementalMigrationStep( + func(tx *sql.Tx) (uint64, error) { + return 10, nil + }, + func(tx *sql.Tx, increment incrementCountFn) error { + // Simulate work with increments + for range 10 { + increment() + time.Sleep(5 * time.Millisecond) + } + return nil + }, + ) + + err = step(tx) + require.NoError(t, err) + + // Verify progress output was written (includes 100% complete at the end) + assert.Len(t, strings.Split(strings.Trim(buf.String(), "\n"), "\n"), 6) + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("increment updates progress", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectBegin() + + tx, err := db.Begin() + require.NoError(t, err) + + // Override output and progress interval for testing + var buf bytes.Buffer + outputTo = &buf + progressInterval = 20 * time.Millisecond + + incrementCount := 0 + step := incrementalMigrationStep( + func(tx *sql.Tx) (uint64, error) { + return 50, nil + }, + func(tx *sql.Tx, increment incrementCountFn) error { + // Call increment multiple times + for range 50 { + increment() + incrementCount++ + } + // Allow time for progress ticker + time.Sleep(30 * time.Millisecond) + return nil + }, + ) + + err = step(tx) + require.NoError(t, err) + assert.Equal(t, 50, incrementCount) + require.Equal(t, " Almost done...\n 100% complete\n", buf.String()) + + require.NoError(t, mock.ExpectationsWereMet()) + }) +} + +func TestWithSteps(t *testing.T) { + // Save original values and restore after test + originalOutputTo := outputTo + defer func() { + outputTo = originalOutputTo + }() + + t.Run("empty steps succeeds", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + var buf bytes.Buffer + outputTo = &buf + + mock.ExpectBegin() + + tx, err := db.Begin() + require.NoError(t, err) + + err = withSteps([]migrationStep{}, tx) + require.NoError(t, err) + + require.NoError(t, mock.ExpectationsWereMet()) + + require.Empty(t, buf.String()) + }) + + t.Run("single step no step output", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectBegin() + + tx, err := db.Begin() + require.NoError(t, err) + + var buf bytes.Buffer + outputTo = &buf + + stepCalled := false + steps := []migrationStep{ + func(tx *sql.Tx) error { + stepCalled = true + return nil + }, + } + + err = withSteps(steps, tx) + require.NoError(t, err) + assert.True(t, stepCalled) + + // Single step should not output step number + output := buf.String() + assert.Empty(t, output) + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("multiple steps with output", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectBegin() + + tx, err := db.Begin() + require.NoError(t, err) + + var buf bytes.Buffer + outputTo = &buf + + var callOrder []int + steps := []migrationStep{ + func(tx *sql.Tx) error { + callOrder = append(callOrder, 1) + return nil + }, + func(tx *sql.Tx) error { + callOrder = append(callOrder, 2) + return nil + }, + func(tx *sql.Tx) error { + callOrder = append(callOrder, 3) + return nil + }, + } + + err = withSteps(steps, tx) + require.NoError(t, err) + assert.Equal(t, []int{1, 2, 3}, callOrder) + + // Multiple steps should output step numbers + assert.Equal(t, " Step 1 of 3\n Step 2 of 3\n Step 3 of 3\n", buf.String()) + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("error stops execution", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectBegin() + + tx, err := db.Begin() + require.NoError(t, err) + + var buf bytes.Buffer + outputTo = &buf + + expectedErr := errors.New("step 2 failed") + var callOrder []int + steps := []migrationStep{ + func(tx *sql.Tx) error { + callOrder = append(callOrder, 1) + return nil + }, + func(tx *sql.Tx) error { + callOrder = append(callOrder, 2) + return expectedErr + }, + func(tx *sql.Tx) error { + callOrder = append(callOrder, 3) + return nil + }, + } + + err = withSteps(steps, tx) + require.Error(t, err) + assert.Equal(t, expectedErr, err) + assert.Equal(t, []int{1, 2}, callOrder, "step 3 should not be called after step 2 fails") + require.Equal(t, " Step 1 of 3\n Step 2 of 3\n", buf.String()) + + require.NoError(t, mock.ExpectationsWereMet()) + }) + + t.Run("integration with basicMigrationStep", func(t *testing.T) { + db, mock, err := sqlmock.New() + require.NoError(t, err) + defer db.Close() + + mock.ExpectBegin() + mock.ExpectExec("ALTER TABLE foo ADD COLUMN a INT").WillReturnResult(sqlmock.NewResult(0, 0)) + mock.ExpectExec("ALTER TABLE foo ADD COLUMN b INT").WillReturnResult(sqlmock.NewResult(0, 0)) + + tx, err := db.Begin() + require.NoError(t, err) + + var buf bytes.Buffer + outputTo = &buf + + steps := []migrationStep{ + basicMigrationStep("ALTER TABLE foo ADD COLUMN a INT", "failed to add column a"), + basicMigrationStep("ALTER TABLE foo ADD COLUMN b INT", "failed to add column b"), + } + + err = withSteps(steps, tx) + require.NoError(t, err) + + require.Equal(t, buf.String(), " Step 1 of 2\n Step 2 of 2\n") + + require.NoError(t, mock.ExpectationsWereMet()) + }) +} + +// TestOutputToAndProgressIntervalDefaults verifies the default values of the package variables +func TestOutputToAndProgressIntervalDefaults(t *testing.T) { + // Note: These tests verify the defaults are sensible + // The actual io.Writer interface check is sufficient + var _ io.Writer = outputTo + assert.Equal(t, 5*time.Second, progressInterval) +}