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
This commit is contained in:
Ian Littman
2026-01-22 15:00:21 -06:00
committed by GitHub
parent 774595f32e
commit 88f8ade624
2 changed files with 483 additions and 4 deletions
@@ -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
@@ -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)
}