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:
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user