diff --git a/cmd/go_migrate/batch-generator.go b/cmd/go_migrate/batch-generator.go deleted file mode 100644 index b0488b9..0000000 --- a/cmd/go_migrate/batch-generator.go +++ /dev/null @@ -1,108 +0,0 @@ -package main - -import ( - "context" - "database/sql" - "fmt" - "time" - - "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" - "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/models" - "github.com/google/uuid" -) - -func estimateTotalRowsMssql(ctx context.Context, db *sql.DB, tableInfo config.SourceTableInfo) (int64, error) { - query := ` -SELECT - SUM(p.rows) AS count -FROM sys.tables t -JOIN sys.schemas s ON t.schema_id = s.schema_id -JOIN sys.partitions p ON t.object_id = p.object_id -WHERE s.name = @schema AND t.name = @table AND p.index_id IN (0, 1) -GROUP BY t.name` - - ctxTimeout, cancel := context.WithTimeout(ctx, time.Second*20) - defer cancel() - - var rowsCount int64 - err := db.QueryRowContext(ctxTimeout, query, sql.Named("schema", tableInfo.Schema), sql.Named("table", tableInfo.Table)).Scan(&rowsCount) - if err != nil { - return 0, err - } - - return rowsCount, nil -} - -func calculateBatchesMssql(ctx context.Context, db *sql.DB, tableInfo config.SourceTableInfo, batchCount int64) ([]models.Partition, error) { - query := fmt.Sprintf(` -SELECT - MIN([%s]) AS lower_limit, - MAX([%s]) AS upper_limit -FROM - (SELECT [%s], NTILE(@batchCount) OVER (ORDER BY [%s]) AS batch_id FROM [%s].[%s]) AS T -GROUP BY batch_id -ORDER BY batch_id`, - tableInfo.PrimaryKey, - tableInfo.PrimaryKey, - tableInfo.PrimaryKey, - tableInfo.PrimaryKey, - tableInfo.Schema, - tableInfo.Table) - - ctxTimeout, cancel := context.WithTimeout(ctx, time.Second*20) - defer cancel() - - rows, err := db.QueryContext(ctxTimeout, query, sql.Named("batchCount", batchCount)) - if err != nil { - return nil, err - } - defer rows.Close() - - batches := make([]models.Partition, 0, batchCount) - - for rows.Next() { - batch := models.Partition{ - Id: uuid.New(), - ShouldUseRange: true, - RetryCounter: 0, - IsLowerLimitInclusive: true, - } - - if err := rows.Scan(&batch.LowerLimit, &batch.UpperLimit); err != nil { - return nil, err - } - - batches = append(batches, batch) - } - - if err := rows.Err(); err != nil { - return nil, err - } - - return batches, nil -} - -func batchGeneratorMssql(ctx context.Context, db *sql.DB, tableInfo config.SourceTableInfo, rowsPerBatch int64) ([]models.Partition, error) { - rowsCount, err := estimateTotalRowsMssql(ctx, db, tableInfo) - if err != nil { - return nil, err - } - - var batchCount int64 = 1 - if rowsCount > rowsPerBatch { - batchCount = rowsCount / rowsPerBatch - } else { - return []models.Partition{{ - Id: uuid.New(), - ShouldUseRange: false, - RetryCounter: 0, - }}, nil - } - - batches, err := calculateBatchesMssql(ctx, db, tableInfo, batchCount) - if err != nil { - return nil, err - } - - return batches, nil -} diff --git a/cmd/go_migrate/colum-type.go b/cmd/go_migrate/colum-type.go deleted file mode 100644 index cfd76a4..0000000 --- a/cmd/go_migrate/colum-type.go +++ /dev/null @@ -1,44 +0,0 @@ -package main - -type ColumnType struct { - name string - - hasMaxLength bool - hasPrecisionScale bool - - userType string - systemType string - unifiedType string - nullable bool - maxLength int64 - precision int64 - scale int64 -} - -func (c *ColumnType) Name() string { - return c.name -} - -func (c *ColumnType) UserType() string { - return c.userType -} - -func (c *ColumnType) SystemType() string { - return c.systemType -} - -func (c *ColumnType) Length() (length int64, ok bool) { - return c.maxLength, c.hasMaxLength -} - -func (c *ColumnType) DecimalSize() (precision, scale int64, ok bool) { - return c.precision, c.scale, c.hasPrecisionScale -} - -func (c *ColumnType) Nullable() bool { - return c.nullable -} - -func (c *ColumnType) Type() string { - return c.unifiedType -} diff --git a/cmd/go_migrate/inspect-columns.go b/cmd/go_migrate/inspect-columns.go deleted file mode 100644 index 6da878c..0000000 --- a/cmd/go_migrate/inspect-columns.go +++ /dev/null @@ -1,316 +0,0 @@ -package main - -import ( - "context" - "database/sql" - "errors" - "fmt" - "strings" - "sync" - "time" - - "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" - "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/models" - "github.com/jackc/pgx/v5/pgxpool" - _ "github.com/microsoft/go-mssqldb" - log "github.com/sirupsen/logrus" -) - -func GetUnifiedType(systemType string) string { - systemType = strings.ToLower(systemType) - - if systemType == "varchar" || systemType == "char" || systemType == "nvarchar" || systemType == "nchar" || systemType == "text" || systemType == "ntext" { - return "STRING" - } - - if systemType == "int" || systemType == "int4" || systemType == "integer" || systemType == "smallint" || systemType == "int2" || systemType == "bigint" || systemType == "int8" || systemType == "tinyint" { - return "INTEGER" - } - - if systemType == "decimal" || systemType == "numeric" { - return "DECIMAL" - } - - if systemType == "float" || systemType == "real" || systemType == "double precision" { - return "FLOAT" - } - - if systemType == "bit" || systemType == "boolean" { - return "BOOLEAN" - } - - if systemType == "date" { - return "DATE" - } - if systemType == "time" || systemType == "time without time zone" { - return "TIME" - } - if systemType == "datetime" || systemType == "datetime2" || systemType == "timestamp" || systemType == "timestamptz" || systemType == "timestamp with time zone" { - return "TIMESTAMP" - } - - if systemType == "binary" || systemType == "varbinary" || systemType == "image" || systemType == "bytea" { - return "BINARY" - } - - if systemType == "uniqueidentifier" || systemType == "uuid" { - return "UUID" - } - - if systemType == "json" { - return "JSON" - } - - if systemType == "geometry" || systemType == "geography" { - return "GEOMETRY" - } - - return strings.ToUpper(systemType) -} - -func MapPostgresColumn(column ColumnType, maxLength *int64, precision *int64, scale *int64) models.ColumnType { - stringTypes := map[string]bool{ - "varchar": true, "char": true, "character": true, "text": true, "character varying": true, - } - - decimalTypes := map[string]bool{ - "decimal": true, "numeric": true, - } - - if stringTypes[column.systemType] { - if maxLength != nil { - column.maxLength = *maxLength - column.hasMaxLength = true - } else { - column.maxLength = -1 - column.hasMaxLength = false - } - column.hasPrecisionScale = false - column.precision = -1 - column.scale = -1 - } else if decimalTypes[column.systemType] { - column.hasMaxLength = false - column.maxLength = -1 - if precision != nil && scale != nil { - column.precision = *precision - column.scale = *scale - column.hasPrecisionScale = true - } else { - column.precision = -1 - column.scale = -1 - column.hasPrecisionScale = false - } - } else { - column.hasMaxLength = false - column.maxLength = -1 - column.hasPrecisionScale = false - column.precision = -1 - column.scale = -1 - } - - column.unifiedType = GetUnifiedType(column.systemType) - - colType := models.NewColumnType( - column.name, - column.hasMaxLength, - column.hasPrecisionScale, - column.userType, - column.systemType, - column.unifiedType, - column.nullable, - column.maxLength, - column.precision, - column.scale, - ) - - return colType -} - -func GetColumnTypesPostgres(db *pgxpool.Pool, tableInfo config.TargetTableInfo) ([]models.ColumnType, error) { - query := ` -SELECT - c.column_name AS name, - c.data_type AS user_type, - c.udt_name AS system_type, - (CASE WHEN c.is_nullable = 'YES' THEN TRUE ELSE FALSE END) AS nullable, - c.character_maximum_length AS max_length, - c.numeric_precision AS precision, - c.numeric_scale AS scale -FROM information_schema.columns c -WHERE c.table_schema = $1 AND c.table_name = $2 -ORDER BY c.ordinal_position; -` - - ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second) - defer cancel() - - rows, err := db.Query(ctx, query, tableInfo.Schema, tableInfo.Table) - if err != nil { - return nil, fmt.Errorf("Error querying column types: %w", err) - } - defer rows.Close() - - var colTypes []models.ColumnType - - for rows.Next() { - var column ColumnType - var scanMaxLength *int64 - var scanPrecision *int64 - var scanScale *int64 - - if err := rows.Scan( - &column.name, - &column.userType, - &column.systemType, - &column.nullable, - &scanMaxLength, - &scanPrecision, - &scanScale, - ); err != nil { - return nil, fmt.Errorf("Error scanning column type results: %w", err) - } - - colTypes = append(colTypes, MapPostgresColumn(column, scanMaxLength, scanPrecision, scanScale)) - } - - return colTypes, nil -} - -func MapMssqlColumn(column ColumnType) models.ColumnType { - stringTypes := map[string]bool{ - "varchar": true, "char": true, "nvarchar": true, "nchar": true, "text": true, "ntext": true, - } - - decimalTypes := map[string]bool{ - "decimal": true, "numeric": true, - } - - if stringTypes[column.systemType] { - column.hasMaxLength = true - if column.systemType == "nvarchar" || column.systemType == "nchar" { - if column.maxLength > 0 { - column.maxLength = column.maxLength / 2 - } - } - column.hasPrecisionScale = false - column.precision = -1 - column.scale = -1 - } else if decimalTypes[column.systemType] { - column.hasMaxLength = false - column.maxLength = -1 - column.hasPrecisionScale = true - } else { - column.hasMaxLength = false - column.maxLength = -1 - column.hasPrecisionScale = false - column.precision = -1 - column.scale = -1 - } - - column.unifiedType = GetUnifiedType(column.systemType) - - colType := models.NewColumnType( - column.name, - column.hasMaxLength, - column.hasPrecisionScale, - column.userType, - column.systemType, - column.unifiedType, - column.nullable, - column.maxLength, - column.precision, - column.scale, - ) - - return colType -} - -func GetColumnTypesMssql(db *sql.DB, tableInfo config.SourceTableInfo) ([]models.ColumnType, error) { - query := ` -SELECT - c.name AS name, - t.name AS user_type, - CASE WHEN t.is_user_defined = 0 THEN t.name ELSE bt.name END AS system_type, - c.is_nullable AS nullable, - c.max_length AS max_length, - c.precision AS precision, - c.scale AS scale -FROM sys.columns c -JOIN sys.types t ON c.user_type_id = t.user_type_id -LEFT JOIN sys.types bt ON t.is_user_defined = 1 AND bt.user_type_id = t.system_type_id -JOIN sys.tables st ON c.object_id = st.object_id -JOIN sys.schemas s ON st.schema_id = s.schema_id -WHERE s.name = @schema AND st.name = @table -ORDER BY c.column_id; -` - - ctx, cancel := context.WithTimeout(context.Background(), 20*time.Second) - defer cancel() - - rows, err := db.QueryContext(ctx, query, sql.Named("schema", tableInfo.Schema), sql.Named("table", tableInfo.Table)) - if err != nil { - return nil, fmt.Errorf("Error querying column types: %w", err) - } - defer rows.Close() - - var colTypes []models.ColumnType - - for rows.Next() { - var column ColumnType - - if err := rows.Scan( - &column.name, - &column.userType, - &column.systemType, - &column.nullable, - &column.maxLength, - &column.precision, - &column.scale, - ); err != nil { - return nil, fmt.Errorf("Error scanning column type results: %W", err) - } - - if strings.HasPrefix(column.name, "graph_id") && column.systemType == "bigint" { - continue - } - - colTypes = append(colTypes, MapMssqlColumn(column)) - } - - return colTypes, nil -} - -func GetColumnTypes( - sourceDb *sql.DB, - targetDb *pgxpool.Pool, - sourceTable config.SourceTableInfo, - targetTable config.TargetTableInfo, -) ([]models.ColumnType, []models.ColumnType, error) { - var sourceDbErr error - var targetDbErr error - var sourceColTypes []models.ColumnType - var targetColTypes []models.ColumnType - var wg sync.WaitGroup - - wg.Go(func() { - sourceColTypes, sourceDbErr = GetColumnTypesMssql(sourceDb, sourceTable) - if sourceDbErr != nil { - log.Error("Error (sourceDb): ", sourceDbErr) - } - }) - - wg.Go(func() { - targetColTypes, targetDbErr = GetColumnTypesPostgres(targetDb, targetTable) - if targetDbErr != nil { - log.Error("Error (targetDb): ", targetDbErr) - } - }) - - wg.Wait() - - if sourceDbErr != nil || targetDbErr != nil { - return nil, nil, errors.New("Error querying column types") - } - - return sourceColTypes, targetColTypes, nil -} diff --git a/cmd/go_migrate/main.go b/cmd/go_migrate/main.go index 6afdf01..5e399fd 100644 --- a/cmd/go_migrate/main.go +++ b/cmd/go_migrate/main.go @@ -9,6 +9,7 @@ import ( "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/etl/extractors" "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/etl/loaders" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/etl/table_analyzers" "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/etl/transformers" "github.com/jackc/pgx/v5/pgxpool" log "github.com/sirupsen/logrus" @@ -90,6 +91,8 @@ func processMigrationJobs( chJobs := make(chan config.Job, len(jobs)) var wgJobs sync.WaitGroup + sourceTableAnalyzer := table_analyzers.NewMssqlTableAnalyzer(sourceDb) + targetTableAnalyzer := table_analyzers.NewPostgresTableAnalyzer(targetDb) extractor := extractors.NewMssqlExtractor(sourceDb) transformer := transformers.NewMssqlTransformer() loader := loaders.NewPostgresLoader(targetDb) @@ -100,8 +103,8 @@ func processMigrationJobs( log.Infof("[worker %d] >>> Processing job: %s.%s <<<", i, job.SourceTable.Schema, job.SourceTable.Table) res := processMigrationJob( ctx, - sourceDb, - targetDb, + sourceTableAnalyzer, + targetTableAnalyzer, extractor, transformer, loader, diff --git a/cmd/go_migrate/process.go b/cmd/go_migrate/process.go index 21b3fc7..3930922 100644 --- a/cmd/go_migrate/process.go +++ b/cmd/go_migrate/process.go @@ -2,7 +2,6 @@ package main import ( "context" - "database/sql" "sync" "sync/atomic" "time" @@ -10,21 +9,24 @@ import ( "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/etl" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/etl/table_analyzers" "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/models" - "github.com/jackc/pgx/v5/pgxpool" - _ "github.com/microsoft/go-mssqldb" log "github.com/sirupsen/logrus" + "golang.org/x/sync/errgroup" ) func processMigrationJob( ctx context.Context, - sourceDb *sql.DB, - targetDb *pgxpool.Pool, + sourceTableAnalyzer etl.TableAnalyzer, + targetTableAnalyzer etl.TableAnalyzer, extractor etl.Extractor, transformer etl.Transformer, loader etl.Loader, job config.Job, ) JobResult { + jobCtx, cancel := context.WithCancel(ctx) + defer cancel() + result := JobResult{ JobName: job.Name, StartTime: time.Now(), @@ -32,47 +34,85 @@ func processMigrationJob( var rowsRead, rowsLoaded, rowsFailed int64 - sourceColTypes, targetColTypes, err := GetColumnTypes(sourceDb, targetDb, job.SourceTable, job.TargetTable) + var wgQueryColumnTypes errgroup.Group + var sourceColTypes, targetColTypes []models.ColumnType + + wgQueryColumnTypes.Go(func() error { + var err error + sourceColTypes, err = sourceTableAnalyzer.QueryColumnTypes(jobCtx, job.SourceTable.TableInfo) + if err != nil { + return err + } + + return nil + }) + + wgQueryColumnTypes.Go(func() error { + var err error + targetColTypes, err = targetTableAnalyzer.QueryColumnTypes(jobCtx, job.TargetTable.TableInfo) + if err != nil { + return err + } + + return nil + }) + + err := wgQueryColumnTypes.Wait() if err != nil { result.Error = err return result } - logColumnTypes(sourceColTypes, "Source col types") - logColumnTypes(targetColTypes, "Target col types") - - jobCtx, cancel := context.WithCancel(ctx) - defer cancel() - - batches, err := batchGeneratorMssql(jobCtx, sourceDb, job.SourceTable, job.RowsPerBatch) + partitions, err := table_analyzers.PartitionRangeGenerator( + jobCtx, + sourceTableAnalyzer, + job.SourceTable.TableInfo, + job.SourceTable.PrimaryKey, + job.RowsPerPartition, + ) if err != nil { log.Error("Unexpected error calculating batch ranges: ", err) } chJobErrors := make(chan custom_errors.JobError, job.QueueSize) - chBatches := make(chan models.Partition, job.QueueSize) chExtractorErrors := make(chan custom_errors.ExtractorError, job.QueueSize) - chChunksRaw := make(chan models.Batch, job.QueueSize) - chChunksTransformed := make(chan models.Batch, job.QueueSize) chLoadersErrors := make(chan custom_errors.LoaderError, job.QueueSize) + chPartitions := make(chan models.Partition, job.QueueSize) + chBatchesRaw := make(chan models.Batch, job.QueueSize) + chBatchesTransformed := make(chan models.Batch, job.QueueSize) + var wgActivePartitions sync.WaitGroup var wgActiveBatches sync.WaitGroup - var wgActiveChunks sync.WaitGroup var wgExtractors sync.WaitGroup var wgTransformers sync.WaitGroup var wgLoaders sync.WaitGroup go func() { if err := custom_errors.JobErrorHandler(jobCtx, chJobErrors); err != nil { + log.Error("Fatal error received from JobErrorHandler, canceling job... - ", err) cancel() result.Error = err } }() - go custom_errors.ExtractorErrorHandler(jobCtx, job.Retry.Attempts, chExtractorErrors, chBatches, chJobErrors, &wgActiveBatches) - go custom_errors.LoaderErrorHandler(jobCtx, job.Retry.Attempts, chLoadersErrors, chChunksTransformed, chJobErrors, &wgActiveChunks) + go custom_errors.ExtractorErrorHandler( + jobCtx, + job.Retry, + chExtractorErrors, + chPartitions, + chJobErrors, + &wgActivePartitions, + ) + go custom_errors.LoaderErrorHandler( + jobCtx, + job.Retry, + chLoadersErrors, + chBatchesTransformed, + chJobErrors, + &wgActiveBatches, + ) - maxExtractors := min(job.MaxExtractors, len(batches)) + maxExtractors := min(job.MaxExtractors, len(partitions)) log.Infof("Starting %d extractor(s)...", maxExtractors) for range maxExtractors { @@ -81,21 +121,21 @@ func processMigrationJob( jobCtx, job.SourceTable, sourceColTypes, - job.ChunkSize, - chBatches, - chChunksRaw, + job.BatchSize, + chPartitions, + chBatchesRaw, chExtractorErrors, chJobErrors, - &wgActiveBatches, + &wgActivePartitions, &rowsRead, ) }) } - wgActiveBatches.Add(len(batches)) + wgActivePartitions.Add(len(partitions)) go func() { - for _, batch := range batches { - chBatches <- batch + for _, batch := range partitions { + chPartitions <- batch } }() @@ -106,10 +146,10 @@ func processMigrationJob( transformer.Exec( jobCtx, sourceColTypes, - chChunksRaw, - chChunksTransformed, + chBatchesRaw, + chBatchesTransformed, chJobErrors, - &wgActiveChunks, + &wgActiveBatches, ) }) } @@ -122,35 +162,49 @@ func processMigrationJob( jobCtx, job.TargetTable, targetColTypes, - chChunksTransformed, + chBatchesTransformed, chLoadersErrors, chJobErrors, - &wgActiveChunks, + &wgActiveBatches, &rowsLoaded, ) }) } go func() { - wgActiveBatches.Wait() - close(chBatches) + log.Debugf("Waiting for goroutines (%v)", job.Name) + + wgActivePartitions.Wait() + log.Debugf("wgActivePartitions is empty (%v)", job.Name) + close(chPartitions) + log.Debugf("chPartitions is closed (%v)", job.Name) close(chExtractorErrors) + log.Debugf("chExtractorErrors is closed (%v)", job.Name) wgExtractors.Wait() - close(chChunksRaw) + log.Debugf("wgExtractors is empty (%v)", job.Name) + close(chBatchesRaw) + log.Debugf("chBatchesRaw is closed (%v)", job.Name) wgTransformers.Wait() + log.Debugf("wgTransformers is empty (%v)", job.Name) - wgActiveChunks.Wait() - close(chChunksTransformed) + wgActiveBatches.Wait() + log.Debugf("wgActiveBatches is empty (%v)", job.Name) + close(chBatchesTransformed) + log.Debugf("chBatchesTransformed is empty (%v)", job.Name) close(chLoadersErrors) + log.Debugf("chLoadersErrors is empty (%v)", job.Name) wgLoaders.Wait() + log.Debugf("wgLoaders is empty (%v)", job.Name) cancel() }() + log.Debugf("waiting for local context to be done (%v)", job.Name) <-jobCtx.Done() + log.Debugf("local context done (%v)", job.Name) if ctx.Err() != nil { result.Error = ctx.Err() @@ -163,11 +217,3 @@ func processMigrationJob( return result } - -func logColumnTypes(columnTypes []models.ColumnType, label string) { - log.Debug(label) - - for _, col := range columnTypes { - log.Debugf("%+v", col) - } -} diff --git a/config.yaml b/config.yaml index ceebcca..caad962 100644 --- a/config.yaml +++ b/config.yaml @@ -6,12 +6,15 @@ defaults: max_extractors: 2 max_loaders: 4 queue_size: 8 - chunk_size: 25000 - chunks_per_batch: 8 + batch_size: 25000 + batches_per_partition: 8 truncate_target: true truncate_method: TRUNCATE # TRUNCATE | DELETE retry: attempts: 3 + base_delay_ms: 500 + max_delay_ms: 10000 + max_jitter_ms: 500 jobs: - name: demo_users diff --git a/go.mod b/go.mod index af6b560..96170be 100644 --- a/go.mod +++ b/go.mod @@ -10,6 +10,7 @@ require ( github.com/microsoft/go-mssqldb v1.9.8 github.com/sirupsen/logrus v1.9.4 github.com/twpayne/go-geom v1.6.1 + golang.org/x/sync v0.19.0 gopkg.in/yaml.v3 v3.0.1 ) @@ -23,7 +24,6 @@ require ( github.com/rogpeppe/go-internal v1.14.1 // indirect github.com/shopspring/decimal v1.4.0 // indirect golang.org/x/crypto v0.48.0 // indirect - golang.org/x/sync v0.19.0 // indirect golang.org/x/sys v0.41.0 // indirect golang.org/x/text v0.34.0 // indirect ) diff --git a/internal/app/config/migration.go b/internal/app/config/migration.go index 9bb6f4c..4819c4a 100644 --- a/internal/app/config/migration.go +++ b/internal/app/config/migration.go @@ -9,18 +9,21 @@ import ( type RetryConfig struct { Attempts int `yaml:"attempts"` + BaseDelayMs int `yaml:"base_delay_ms"` + MaxDelayMs int `yaml:"max_delay_ms"` + MaxJitterMs int `yaml:"max_jitter_ms"` } type JobConfig struct { - MaxExtractors int `yaml:"max_extractors"` - MaxLoaders int `yaml:"max_loaders"` - QueueSize int `yaml:"queue_size"` - ChunkSize int `yaml:"chunk_size"` - ChunksPerBatch int `yaml:"chunks_per_batch"` - RowsPerBatch int64 - TruncateTarget bool `yaml:"truncate_target"` - TruncateMethod string `yaml:"truncate_method"` - Retry RetryConfig `yaml:"retry"` + MaxExtractors int `yaml:"max_extractors"` + MaxLoaders int `yaml:"max_loaders"` + QueueSize int `yaml:"queue_size"` + BatchSize int `yaml:"batch_size"` + BatchesPerPartition int `yaml:"batches_per_partition"` + TruncateTarget bool `yaml:"truncate_target"` + TruncateMethod string `yaml:"truncate_method"` + Retry RetryConfig `yaml:"retry"` + RowsPerPartition int64 } type TableInfo struct { @@ -71,7 +74,7 @@ func (c *MigrationConfig) UnmarshalYAML(value *yaml.Node) error { c.MaxParallelWorkers = raw.MaxParallelWorkers c.Defaults = raw.Defaults - c.Defaults.RowsPerBatch = int64(raw.Defaults.ChunkSize * raw.Defaults.ChunksPerBatch) + c.Defaults.RowsPerPartition = int64(raw.Defaults.BatchSize * raw.Defaults.BatchesPerPartition) for _, node := range raw.Jobs { job := Job{ @@ -82,7 +85,7 @@ func (c *MigrationConfig) UnmarshalYAML(value *yaml.Node) error { return err } - job.RowsPerBatch = int64(job.ChunkSize * job.ChunksPerBatch) + job.RowsPerPartition = int64(job.BatchSize * job.BatchesPerPartition) c.Jobs = append(c.Jobs, job) } diff --git a/internal/app/custom_errors/backoff.go b/internal/app/custom_errors/backoff.go new file mode 100644 index 0000000..fc469da --- /dev/null +++ b/internal/app/custom_errors/backoff.go @@ -0,0 +1,61 @@ +package custom_errors + +import ( + "context" + "math/rand" + "time" +) + +func computeBackoffDelay(retryCounter int, baseDelayMs int, maxDelayMs int, maxJitterMs int) time.Duration { + if retryCounter < 0 { + retryCounter = 0 + } + + delay := max(time.Duration(baseDelayMs)*time.Millisecond, 0) + + maxDelay := time.Duration(maxDelayMs) * time.Millisecond + for i := 0; i < retryCounter; i++ { + if maxDelayMs > 0 && delay >= maxDelay { + delay = maxDelay + break + } + if delay == 0 { + break + } + delay *= 2 + } + + if maxDelayMs > 0 && delay > maxDelay { + delay = maxDelay + } + + if maxJitterMs > 0 { + jitter := time.Duration(rand.Intn(maxJitterMs+1)) * time.Millisecond + delay += jitter + } + + if delay < 0 { + delay = 0 + } + + return delay +} + +func requeueWithBackoff(ctx context.Context, delay time.Duration, enqueue func()) { + if delay <= 0 { + enqueue() + return + } + + go func() { + timer := time.NewTimer(delay) + defer timer.Stop() + + select { + case <-ctx.Done(): + return + case <-timer.C: + enqueue() + } + }() +} diff --git a/internal/app/custom_errors/extractor.error.go b/internal/app/custom_errors/extractor.error.go index ad4ccf0..5ab02e7 100644 --- a/internal/app/custom_errors/extractor.error.go +++ b/internal/app/custom_errors/extractor.error.go @@ -5,12 +5,13 @@ import ( "fmt" "sync" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/models" "github.com/google/uuid" ) type ExtractorError struct { - Batch models.Partition + Partition models.Partition LastId int64 HasLastId bool Msg string @@ -22,11 +23,11 @@ func (e *ExtractorError) Error() string { func ExtractorErrorHandler( ctx context.Context, - maxRetryAttempts int, + retryConfig config.RetryConfig, chErrorsIn <-chan ExtractorError, - chBatchesOut chan<- models.Partition, + chPartitionsOut chan<- models.Partition, chJobErrorsOut chan<- JobError, - wgActiveBatches *sync.WaitGroup, + wgActivePartitions *sync.WaitGroup, ) { for { if ctx.Err() != nil { @@ -42,10 +43,11 @@ func ExtractorErrorHandler( return } - if err.Batch.RetryCounter >= maxRetryAttempts { + if err.Partition.RetryCounter >= retryConfig.Attempts { + wgActivePartitions.Done() jobError := JobError{ ShouldCancelJob: false, - Msg: fmt.Sprintf("batch %v reached max retries (%d)", err.Batch.Id, maxRetryAttempts), + Msg: fmt.Sprintf("Partition %v reached max retries (%d)", err.Partition.Id, retryConfig.Attempts), Prev: &err, } @@ -55,25 +57,45 @@ func ExtractorErrorHandler( return } - wgActiveBatches.Done() continue + } else { + jobError := JobError{ + ShouldCancelJob: false, + Msg: fmt.Sprintf("Temporal error in partition %v (retries: %d)", err.Partition.Id, err.Partition.RetryCounter), + Prev: &err, + } + + select { + case chJobErrorsOut <- jobError: + case <-ctx.Done(): + return + } } - newBatch := err.Batch - newBatch.RetryCounter++ + newPartition := err.Partition + newPartition.RetryCounter++ + + delay := computeBackoffDelay( + newPartition.RetryCounter, + retryConfig.BaseDelayMs, + retryConfig.MaxDelayMs, + retryConfig.MaxJitterMs, + ) if err.HasLastId { - newBatch.ParentId = err.Batch.Id - newBatch.Id = uuid.New() - newBatch.LowerLimit = err.LastId - newBatch.IsLowerLimitInclusive = false + newPartition.ParentId = err.Partition.Id + newPartition.Id = uuid.New() + newPartition.LowerLimit = err.LastId + newPartition.IsLowerLimitInclusive = false } - select { - case chBatchesOut <- newBatch: - case <-ctx.Done(): - return - } + requeueWithBackoff(ctx, delay, func() { + select { + case chPartitionsOut <- newPartition: + case <-ctx.Done(): + return + } + }) } } } diff --git a/internal/app/custom_errors/job.error.go b/internal/app/custom_errors/job.error.go index ca359af..8326068 100644 --- a/internal/app/custom_errors/job.error.go +++ b/internal/app/custom_errors/job.error.go @@ -37,11 +37,11 @@ func JobErrorHandler(ctx context.Context, chErrorsIn <-chan JobError) error { } if err.ShouldCancelJob { - log.Error(err.Msg, " - ", err.Prev) + log.Errorf("(Fatal job error) - %v - %v", err.Msg, err.Prev) return &err } - log.Error(err.Msg, " - ", err.Prev) + log.Errorf("%v - %v", err.Msg, err.Prev) } } } diff --git a/internal/app/custom_errors/loader.error.go b/internal/app/custom_errors/loader.error.go index 4f03e9d..30dae24 100644 --- a/internal/app/custom_errors/loader.error.go +++ b/internal/app/custom_errors/loader.error.go @@ -5,12 +5,13 @@ import ( "fmt" "sync" + "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/config" "git.ksdemosapps.com/kylesoda/go-migrate/internal/app/models" ) type LoaderError struct { - models.Batch - Msg string + Batch models.Batch + Msg string } func (e *LoaderError) Error() string { @@ -19,11 +20,11 @@ func (e *LoaderError) Error() string { func LoaderErrorHandler( ctx context.Context, - maxRetryAttempts int, + retryConfig config.RetryConfig, chErrorsIn <-chan LoaderError, - chChunksOut chan<- models.Batch, + chBatchesOut chan<- models.Batch, chJobErrorsOut chan<- JobError, - wgActiveChunks *sync.WaitGroup, + wgActiveBatches *sync.WaitGroup, ) { for { if ctx.Err() != nil { @@ -39,10 +40,11 @@ func LoaderErrorHandler( return } - if err.RetryCounter >= maxRetryAttempts { + if err.Batch.RetryCounter >= retryConfig.Attempts { + wgActiveBatches.Done() jobError := JobError{ ShouldCancelJob: false, - Msg: fmt.Sprintf("chunk %v reached max retries (%d)", err.Id, maxRetryAttempts), + Msg: fmt.Sprintf("Batch %v reached max retries (%d)", err.Batch.Id, retryConfig.Attempts), Prev: &err, } @@ -52,17 +54,36 @@ func LoaderErrorHandler( return } - wgActiveChunks.Done() continue + } else { + jobError := JobError{ + ShouldCancelJob: false, + Msg: fmt.Sprintf("Temporal error in batch %v (retries: %d)", err.Batch.Id, err.Batch.RetryCounter), + Prev: &err, + } + + select { + case chJobErrorsOut <- jobError: + case <-ctx.Done(): + return + } } - err.RetryCounter++ + err.Batch.RetryCounter++ + delay := computeBackoffDelay( + err.Batch.RetryCounter, + retryConfig.BaseDelayMs, + retryConfig.MaxDelayMs, + retryConfig.MaxJitterMs, + ) - select { - case chChunksOut <- err.Batch: - case <-ctx.Done(): - return - } + requeueWithBackoff(ctx, delay, func() { + select { + case chBatchesOut <- err.Batch: + case <-ctx.Done(): + return + } + }) } } } diff --git a/internal/app/etl/extractors/mssql.go b/internal/app/etl/extractors/mssql.go index 629377d..40e9c0c 100644 --- a/internal/app/etl/extractors/mssql.go +++ b/internal/app/etl/extractors/mssql.go @@ -70,20 +70,20 @@ func buildExtractQueryMssql( return sbQuery.String() } -func extractorErrorFromLastRowMssql( +func errorFromLastRow( lastRow models.UnknownRowValues, indexPrimaryKey int, - batch *models.Partition, + partition *models.Partition, previousError error, ) *custom_errors.ExtractorError { lastIdRawValue := lastRow[indexPrimaryKey] lastId, ok := convert.ToInt64(lastIdRawValue) if !ok { - currentBatch := *batch - currentBatch.RetryCounter = 3 + currentPartition := *partition + currentPartition.RetryCounter = 3 return &custom_errors.ExtractorError{ - Batch: currentBatch, + Partition: currentPartition, HasLastId: true, Msg: fmt.Sprintf("Couldn't cast last id value as int: %s", previousError.Error()), } @@ -91,78 +91,78 @@ func extractorErrorFromLastRowMssql( } return &custom_errors.ExtractorError{ - Batch: *batch, + Partition: *partition, HasLastId: true, LastId: lastId, Msg: previousError.Error(), } } -func (mssqlEx *MssqlExtractor) ProcessBatch( +func (mssqlEx *MssqlExtractor) ProcessPartition( ctx context.Context, tableInfo config.SourceTableInfo, columns []models.ColumnType, - chunkSize int, - batch models.Partition, + batchSize int, + partition models.Partition, indexPrimaryKey int, - chChunksOut chan<- models.Batch, + chBatchesOut chan<- models.Batch, rowsRead *int64, ) error { - query := buildExtractQueryMssql(tableInfo, columns, batch.ShouldUseRange, batch.IsLowerLimitInclusive) + query := buildExtractQueryMssql(tableInfo, columns, partition.ShouldUseRange, partition.IsLowerLimitInclusive) var queryArgs []any - if batch.ShouldUseRange { + if partition.ShouldUseRange { queryArgs = append(queryArgs, - sql.Named("min", batch.LowerLimit), - sql.Named("max", batch.UpperLimit), + sql.Named("min", partition.LowerLimit), + sql.Named("max", partition.UpperLimit), ) } rows, err := mssqlEx.db.QueryContext(ctx, query, queryArgs...) if err != nil { - return &custom_errors.ExtractorError{Batch: batch, HasLastId: false, Msg: err.Error()} + return &custom_errors.ExtractorError{Partition: partition, HasLastId: false, Msg: err.Error()} } defer rows.Close() - rowsChunk := make([]models.UnknownRowValues, 0, chunkSize) + batchRows := make([]models.UnknownRowValues, 0, batchSize) for rows.Next() { - values := make([]any, len(columns)) + rowValues := make([]any, len(columns)) scanArgs := make([]any, len(columns)) - for i := range values { - scanArgs[i] = &values[i] + for i := range rowValues { + scanArgs[i] = &rowValues[i] } if err := rows.Scan(scanArgs...); err != nil { - if len(rowsChunk) == 0 { - return &custom_errors.ExtractorError{Batch: batch, HasLastId: false, Msg: err.Error()} + if len(batchRows) == 0 { + return &custom_errors.ExtractorError{Partition: partition, HasLastId: false, Msg: err.Error()} } - lastRow := rowsChunk[len(rowsChunk)-1] + lastRow := batchRows[len(batchRows)-1] select { - case chChunksOut <- models.Batch{Id: uuid.New(), PartitionId: batch.Id, Data: rowsChunk, RetryCounter: 0}: + case chBatchesOut <- models.Batch{Id: uuid.New(), PartitionId: partition.Id, Rows: batchRows, RetryCounter: 0}: case <-ctx.Done(): return nil } - atomic.AddInt64(rowsRead, int64(len(rowsChunk))) + atomic.AddInt64(rowsRead, int64(len(batchRows))) - return extractorErrorFromLastRowMssql(lastRow, indexPrimaryKey, &batch, err) + return errorFromLastRow(lastRow, indexPrimaryKey, &partition, err) } - rowsChunk = append(rowsChunk, values) + batchRows = append(batchRows, rowValues) - if len(rowsChunk) >= chunkSize { + if len(batchRows) >= batchSize { select { - case chChunksOut <- models.Batch{Id: uuid.New(), PartitionId: batch.Id, Data: rowsChunk, RetryCounter: 0}: + case chBatchesOut <- models.Batch{Id: uuid.New(), PartitionId: partition.Id, Rows: batchRows, RetryCounter: 0}: case <-ctx.Done(): return nil } - atomic.AddInt64(rowsRead, int64(len(rowsChunk))) - rowsChunk = make([]models.UnknownRowValues, 0, chunkSize) + atomic.AddInt64(rowsRead, int64(len(batchRows))) + batchRows = make([]models.UnknownRowValues, 0, batchSize) } } @@ -171,22 +171,22 @@ func (mssqlEx *MssqlExtractor) ProcessBatch( return ctx.Err() } - if len(rowsChunk) == 0 { - return &custom_errors.ExtractorError{Batch: batch, HasLastId: false, Msg: err.Error()} + if len(batchRows) == 0 { + return &custom_errors.ExtractorError{Partition: partition, HasLastId: false, Msg: err.Error()} } - lastRow := rowsChunk[len(rowsChunk)-1] - return extractorErrorFromLastRowMssql(lastRow, indexPrimaryKey, &batch, err) + lastRow := batchRows[len(batchRows)-1] + return errorFromLastRow(lastRow, indexPrimaryKey, &partition, err) } - if len(rowsChunk) > 0 { + if len(batchRows) > 0 { select { - case chChunksOut <- models.Batch{Id: uuid.New(), PartitionId: batch.Id, Data: rowsChunk, RetryCounter: 0}: + case chBatchesOut <- models.Batch{Id: uuid.New(), PartitionId: partition.Id, Rows: batchRows, RetryCounter: 0}: case <-ctx.Done(): return nil } - atomic.AddInt64(rowsRead, int64(len(rowsChunk))) + atomic.AddInt64(rowsRead, int64(len(batchRows))) } return nil @@ -196,12 +196,12 @@ func (mssqlEx *MssqlExtractor) Exec( ctx context.Context, tableInfo config.SourceTableInfo, columns []models.ColumnType, - chunkSize int, - chBatchesIn <-chan models.Partition, - chChunksOut chan<- models.Batch, + batchSize int, + chPartitionsIn <-chan models.Partition, + chBatchesOut chan<- models.Batch, chErrorsOut chan<- custom_errors.ExtractorError, chJobErrorsOut chan<- custom_errors.JobError, - wgActiveBatches *sync.WaitGroup, + wgActivePartitions *sync.WaitGroup, rowsRead *int64, ) { indexPrimaryKey := slices.IndexFunc(columns, func(col models.ColumnType) bool { @@ -229,45 +229,49 @@ func (mssqlEx *MssqlExtractor) Exec( select { case <-ctx.Done(): return - case batch, ok := <-chBatchesIn: + case partition, ok := <-chPartitionsIn: if !ok { return } - err := mssqlEx.ProcessBatch( + err := mssqlEx.ProcessPartition( ctx, tableInfo, columns, - chunkSize, - batch, + batchSize, + partition, indexPrimaryKey, - chChunksOut, + chBatchesOut, rowsRead, ) if err != nil { var exError *custom_errors.ExtractorError + var jobError *custom_errors.JobError if errors.As(err, &exError) { select { case <-ctx.Done(): return case chErrorsOut <- *exError: } - } - - var jobError *custom_errors.JobError - if errors.As(err, &jobError) { + } else if errors.As(err, &jobError) { select { case <-ctx.Done(): return case chJobErrorsOut <- *jobError: } + } else { + select { + case <-ctx.Done(): + return + case chErrorsOut <- custom_errors.ExtractorError{Partition: partition, Msg: err.Error()}: + } } - return + continue } - wgActiveBatches.Done() + wgActivePartitions.Done() } } } diff --git a/internal/app/etl/extractors/postgres.go b/internal/app/etl/extractors/postgres.go index 374b5da..6cd1d3a 100644 --- a/internal/app/etl/extractors/postgres.go +++ b/internal/app/etl/extractors/postgres.go @@ -52,29 +52,29 @@ func buildExtractQueryPostgres(sourceDbInfo config.SourceTableInfo, columns []mo return fmt.Sprintf(`SELECT %s FROM "%s"."%s" ORDER BY "%s" ASC`, sbColumns.String(), sourceDbInfo.Schema, sourceDbInfo.Table, sourceDbInfo.PrimaryKey) } -func (postgresEx *PostgresExtractor) ProcessBatch( +func (postgresEx *PostgresExtractor) ProcessPartition( ctx context.Context, tableInfo config.SourceTableInfo, columns []models.ColumnType, - chunkSize int, - batch models.Partition, + batchSize int, + partition models.Partition, indexPrimaryKey int, - chChunksOut chan<- models.Batch, + chBatchesOut chan<- models.Batch, rowsRead *int64, ) error { query := buildExtractQueryPostgres(tableInfo, columns) - if batch.ShouldUseRange { + if partition.ShouldUseRange { return errors.New("Batch config not yet supported") } rows, err := postgresEx.db.Query(ctx, query) if err != nil { - return &custom_errors.ExtractorError{Batch: batch, HasLastId: false, Msg: err.Error()} + return &custom_errors.ExtractorError{Partition: partition, HasLastId: false, Msg: err.Error()} } defer rows.Close() - rowsChunk := make([]models.UnknownRowValues, 0, chunkSize) + batchRows := make([]models.UnknownRowValues, 0, batchSize) for rows.Next() { values, err := rows.Values() @@ -82,17 +82,17 @@ func (postgresEx *PostgresExtractor) ProcessBatch( return errors.New("Unexpected error reading rows from source") } - rowsChunk = append(rowsChunk, values) + batchRows = append(batchRows, values) - if len(rowsChunk) >= chunkSize { + if len(batchRows) >= batchSize { select { - case chChunksOut <- models.Batch{Id: uuid.New(), PartitionId: batch.Id, Data: rowsChunk, RetryCounter: 0}: + case chBatchesOut <- models.Batch{Id: uuid.New(), PartitionId: partition.Id, Rows: batchRows, RetryCounter: 0}: case <-ctx.Done(): return nil } - atomic.AddInt64(rowsRead, int64(len(rowsChunk))) - rowsChunk = make([]models.UnknownRowValues, 0, chunkSize) + atomic.AddInt64(rowsRead, int64(len(batchRows))) + batchRows = make([]models.UnknownRowValues, 0, batchSize) } } @@ -100,14 +100,14 @@ func (postgresEx *PostgresExtractor) ProcessBatch( return errors.New("Unexpected error reading rows from source") } - if len(rowsChunk) > 0 { + if len(batchRows) > 0 { select { - case chChunksOut <- models.Batch{Id: uuid.New(), PartitionId: batch.Id, Data: rowsChunk, RetryCounter: 0}: + case chBatchesOut <- models.Batch{Id: uuid.New(), PartitionId: partition.Id, Rows: batchRows, RetryCounter: 0}: case <-ctx.Done(): return nil } - atomic.AddInt64(rowsRead, int64(len(rowsChunk))) + atomic.AddInt64(rowsRead, int64(len(batchRows))) } return nil @@ -117,12 +117,12 @@ func (postgresEx *PostgresExtractor) Exec( ctx context.Context, tableInfo config.SourceTableInfo, columns []models.ColumnType, - chunkSize int, - chBatchesIn <-chan models.Partition, - chChunksOut chan<- models.Batch, + batchSize int, + chPartitionsIn <-chan models.Partition, + chBatchesOut chan<- models.Batch, chErrorsOut chan<- custom_errors.ExtractorError, chJobErrorsOut chan<- custom_errors.JobError, - wgActiveBatches *sync.WaitGroup, + wgActivePartitions *sync.WaitGroup, rowsRead *int64, ) { } diff --git a/internal/app/etl/loaders/postgres.go b/internal/app/etl/loaders/postgres.go index c9d2276..4560a7c 100644 --- a/internal/app/etl/loaders/postgres.go +++ b/internal/app/etl/loaders/postgres.go @@ -34,18 +34,18 @@ func mapSlice[T any, V any](input []T, mapper func(T) V) []V { return result } -func (postgresLd *PostgresLoader) ProcessChunk( +func (postgresLd *PostgresLoader) ProcessBatch( ctx context.Context, tableInfo config.TargetTableInfo, colNames []string, - chunk models.Batch, + batch models.Batch, ) (int, error) { tableId := pgx.Identifier{tableInfo.Schema, tableInfo.Table} _, err := postgresLd.db.CopyFrom( ctx, tableId, colNames, - pgx.CopyFromRows(chunk.Data), + pgx.CopyFromRows(batch.Rows), ) if err != nil { @@ -60,20 +60,20 @@ func (postgresLd *PostgresLoader) ProcessChunk( } } - return 0, &custom_errors.LoaderError{Batch: chunk, Msg: err.Error()} + return 0, &custom_errors.LoaderError{Batch: batch, Msg: err.Error()} } - return len(chunk.Data), nil + return len(batch.Rows), nil } func (postgresLd *PostgresLoader) Exec( ctx context.Context, tableInfo config.TargetTableInfo, columns []models.ColumnType, - chChunksIn <-chan models.Batch, + chBatchesIn <-chan models.Batch, chErrorsOut chan<- custom_errors.LoaderError, chJobErrorsOut chan<- custom_errors.JobError, - wgActiveChunks *sync.WaitGroup, + wgActiveBatches *sync.WaitGroup, rowsLoaded *int64, ) { colNames := mapSlice(columns, func(col models.ColumnType) string { @@ -88,36 +88,40 @@ func (postgresLd *PostgresLoader) Exec( select { case <-ctx.Done(): return - case chunk, ok := <-chChunksIn: + case batch, ok := <-chBatchesIn: if !ok { return } - processedRows, err := postgresLd.ProcessChunk(ctx, tableInfo, colNames, chunk) + processedRows, err := postgresLd.ProcessBatch(ctx, tableInfo, colNames, batch) if err != nil { var ldError *custom_errors.LoaderError + var jobError *custom_errors.JobError if errors.As(err, &ldError) { select { case <-ctx.Done(): return case chErrorsOut <- *ldError: } - } - - var jobError *custom_errors.JobError - if errors.As(err, &jobError) { + } else if errors.As(err, &jobError) { select { case <-ctx.Done(): return case chJobErrorsOut <- *jobError: } + } else { + select { + case <-ctx.Done(): + return + case chErrorsOut <- custom_errors.LoaderError{Batch: batch, Msg: err.Error()}: + } } - return + continue } - wgActiveChunks.Done() + wgActiveBatches.Done() atomic.AddInt64(rowsLoaded, int64(processedRows)) } } diff --git a/internal/app/etl/table_analyzers/main.go b/internal/app/etl/table_analyzers/main.go new file mode 100644 index 0000000..1ed61d0 --- /dev/null +++ b/internal/app/etl/table_analyzers/main.go @@ -0,0 +1,40 @@ +package table_analyzers + +import ( + "context" + + "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" + "github.com/google/uuid" +) + +func PartitionRangeGenerator( + ctx context.Context, + tableAnalyzer etl.TableAnalyzer, + tableInfo config.TableInfo, + partitionColumn string, + rowsPerPartition int64, +) ([]models.Partition, error) { + rowsCount, err := tableAnalyzer.EstimateTotalRows(ctx, tableInfo) + if err != nil { + return nil, err + } + + if rowsCount <= rowsPerPartition { + return []models.Partition{{ + Id: uuid.New(), + ShouldUseRange: false, + RetryCounter: 0, + }}, nil + + } + + partitionsCount := rowsCount / rowsPerPartition + partitions, err := tableAnalyzer.CalculatePartitionRanges(ctx, tableInfo, partitionColumn, partitionsCount) + if err != nil { + return nil, err + } + + return partitions, nil +} diff --git a/internal/app/etl/table_analyzers/mssql.go b/internal/app/etl/table_analyzers/mssql.go new file mode 100644 index 0000000..abe6ac6 --- /dev/null +++ b/internal/app/etl/table_analyzers/mssql.go @@ -0,0 +1,249 @@ +package table_analyzers + +import ( + "context" + "database/sql" + "fmt" + "strings" + "time" + + "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" + "github.com/google/uuid" +) + +type MssqlTableAnalyzer struct { + db *sql.DB +} + +func NewMssqlTableAnalyzer(db *sql.DB) etl.TableAnalyzer { + return &MssqlTableAnalyzer{db: db} +} + +const mssqlColumnMetadataQuery string = ` +SELECT + c.name AS name, + t.name AS user_type, + CASE WHEN t.is_user_defined = 0 THEN t.name ELSE bt.name END AS system_type, + c.is_nullable AS nullable, + c.max_length AS max_length, + c.precision AS precision, + c.scale AS scale +FROM sys.columns c +JOIN sys.types t ON c.user_type_id = t.user_type_id +LEFT JOIN sys.types bt ON t.is_user_defined = 1 AND bt.user_type_id = t.system_type_id +JOIN sys.tables st ON c.object_id = st.object_id +JOIN sys.schemas s ON st.schema_id = s.schema_id +WHERE s.name = @schema AND st.name = @table AND c.name NOT LIKE 'graph_id%' +ORDER BY c.column_id;` + +type rawColumnMssql struct { + name string + userType string + systemType string + nullable bool + maxLength int64 + precision int64 + scale int64 +} + +func (ta *MssqlTableAnalyzer) systemTypeToUnifiedType(systemType string) string { + systemType = strings.ToLower(systemType) + + if systemType == "varchar" || systemType == "char" || systemType == "nvarchar" || systemType == "nchar" || systemType == "text" || systemType == "ntext" { + return "STRING" + } + + if systemType == "int" || systemType == "int4" || systemType == "integer" || systemType == "smallint" || systemType == "int2" || systemType == "bigint" || systemType == "int8" || systemType == "tinyint" { + return "INTEGER" + } + + if systemType == "decimal" || systemType == "numeric" { + return "DECIMAL" + } + + if systemType == "float" || systemType == "real" || systemType == "double precision" { + return "FLOAT" + } + + if systemType == "bit" || systemType == "boolean" { + return "BOOLEAN" + } + + if systemType == "date" { + return "DATE" + } + if systemType == "time" || systemType == "time without time zone" { + return "TIME" + } + if systemType == "datetime" || systemType == "datetime2" || systemType == "timestamp" || systemType == "timestamptz" || systemType == "timestamp with time zone" { + return "TIMESTAMP" + } + + if systemType == "binary" || systemType == "varbinary" || systemType == "image" || systemType == "bytea" { + return "BINARY" + } + + if systemType == "uniqueidentifier" || systemType == "uuid" { + return "UUID" + } + + if systemType == "json" { + return "JSON" + } + + if systemType == "geometry" || systemType == "geography" { + return "GEOMETRY" + } + + return strings.ToUpper(systemType) +} + +func (ta *MssqlTableAnalyzer) rawColumnToColumnType(rawColumn rawColumnMssql) models.ColumnType { + const nullValue int64 = -1 + stringTypes := map[string]bool{"varchar": true, "char": true, "nvarchar": true, "nchar": true, "text": true, "ntext": true} + decimalTypes := map[string]bool{"decimal": true, "numeric": true} + + if stringTypes[rawColumn.systemType] { + if rawColumn.systemType == "nvarchar" || rawColumn.systemType == "nchar" { + if rawColumn.maxLength > 0 { + rawColumn.maxLength = rawColumn.maxLength / 2 + } + } + + rawColumn.precision, rawColumn.scale = nullValue, nullValue + } else if decimalTypes[rawColumn.systemType] { + rawColumn.maxLength = nullValue + } else { + rawColumn.maxLength, rawColumn.precision, rawColumn.scale = nullValue, nullValue, nullValue + } + + columnType := models.NewColumnType( + rawColumn.name, + rawColumn.maxLength != nullValue, + rawColumn.precision != nullValue || rawColumn.scale != nullValue, + rawColumn.userType, + rawColumn.systemType, + ta.systemTypeToUnifiedType(rawColumn.systemType), + rawColumn.nullable, + rawColumn.maxLength, + rawColumn.precision, + rawColumn.scale, + ) + + return columnType +} + +func (ta *MssqlTableAnalyzer) QueryColumnTypes( + ctx context.Context, + tableInfo config.TableInfo, +) ([]models.ColumnType, error) { + localCtx, cancel := context.WithTimeout(ctx, 20*time.Second) + defer cancel() + + rows, err := ta.db.QueryContext(localCtx, mssqlColumnMetadataQuery, sql.Named("schema", tableInfo.Schema), sql.Named("table", tableInfo.Table)) + if err != nil { + return nil, err + } + defer rows.Close() + + var columnTypes []models.ColumnType + + for rows.Next() { + var rawColumn rawColumnMssql + + if err := rows.Scan( + &rawColumn.name, + &rawColumn.userType, + &rawColumn.systemType, + &rawColumn.nullable, + &rawColumn.maxLength, + &rawColumn.precision, + &rawColumn.scale, + ); err != nil { + return nil, err + } + + columnTypes = append(columnTypes, ta.rawColumnToColumnType(rawColumn)) + } + + return columnTypes, nil +} + +func (ta *MssqlTableAnalyzer) EstimateTotalRows( + ctx context.Context, + tableInfo config.TableInfo, +) (int64, error) { + query := ` +SELECT SUM(p.rows) AS count +FROM sys.tables t +JOIN sys.schemas s ON t.schema_id = s.schema_id +JOIN sys.partitions p ON t.object_id = p.object_id +WHERE s.name = @schema AND t.name = @table AND p.index_id IN (0, 1) +GROUP BY t.name` + + ctxTimeout, cancel := context.WithTimeout(ctx, time.Second*20) + defer cancel() + + var rowsCount int64 + err := ta.db.QueryRowContext(ctxTimeout, query, sql.Named("schema", tableInfo.Schema), sql.Named("table", tableInfo.Table)).Scan(&rowsCount) + if err != nil { + return 0, err + } + + return rowsCount, nil +} + +func (ta *MssqlTableAnalyzer) CalculatePartitionRanges( + ctx context.Context, + tableInfo config.TableInfo, + partitionColumn string, + maxPartitions int64, +) ([]models.Partition, error) { + 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 +GROUP BY batch_id +ORDER BY batch_id`, + partitionColumn, + partitionColumn, + partitionColumn, + partitionColumn, + tableInfo.Schema, + tableInfo.Table) + + ctxTimeout, cancel := context.WithTimeout(ctx, time.Second*20) + defer cancel() + + rows, err := ta.db.QueryContext(ctxTimeout, query, sql.Named("maxPartitions", maxPartitions)) + if err != nil { + return nil, err + } + defer rows.Close() + + partitions := make([]models.Partition, 0, maxPartitions) + + for rows.Next() { + partition := models.Partition{ + Id: uuid.New(), + ShouldUseRange: true, + RetryCounter: 0, + IsLowerLimitInclusive: true, + } + + if err := rows.Scan(&partition.LowerLimit, &partition.UpperLimit); err != nil { + return nil, err + } + + partitions = append(partitions, partition) + } + + if err := rows.Err(); err != nil { + return nil, err + } + + return partitions, nil +} diff --git a/internal/app/etl/table_analyzers/postgres.go b/internal/app/etl/table_analyzers/postgres.go new file mode 100644 index 0000000..b07eb5e --- /dev/null +++ b/internal/app/etl/table_analyzers/postgres.go @@ -0,0 +1,174 @@ +package table_analyzers + +import ( + "context" + "strings" + "time" + + "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" + "github.com/jackc/pgx/v5/pgxpool" +) + +type PostgresTableAnalyzer struct { + db *pgxpool.Pool +} + +func NewPostgresTableAnalyzer(db *pgxpool.Pool) etl.TableAnalyzer { + return &PostgresTableAnalyzer{db: db} +} + +const postgresColumnMetadataQuery string = ` +SELECT + c.column_name AS name, + c.data_type AS user_type, + c.udt_name AS system_type, + (CASE WHEN c.is_nullable = 'YES' THEN TRUE ELSE FALSE END) AS nullable, + COALESCE(c.character_maximum_length, -1) AS max_length, + COALESCE(c.numeric_precision, -1) AS precision, + COALESCE(c.numeric_scale, -1) AS scale +FROM information_schema.columns c +WHERE c.table_schema = $1 AND c.table_name = $2 +ORDER BY c.ordinal_position;` + +type rawColumnPostgres struct { + name string + userType string + systemType string + nullable bool + maxLength int64 + precision int64 + scale int64 +} + +func (ta *PostgresTableAnalyzer) systemTypeToUnifiedType(systemType string) string { + systemType = strings.ToLower(systemType) + + if systemType == "varchar" || systemType == "char" || systemType == "nvarchar" || systemType == "nchar" || systemType == "text" || systemType == "ntext" { + return "STRING" + } + + if systemType == "int" || systemType == "int4" || systemType == "integer" || systemType == "smallint" || systemType == "int2" || systemType == "bigint" || systemType == "int8" || systemType == "tinyint" { + return "INTEGER" + } + + if systemType == "decimal" || systemType == "numeric" { + return "DECIMAL" + } + + if systemType == "float" || systemType == "real" || systemType == "double precision" { + return "FLOAT" + } + + if systemType == "bit" || systemType == "boolean" { + return "BOOLEAN" + } + + if systemType == "date" { + return "DATE" + } + if systemType == "time" || systemType == "time without time zone" { + return "TIME" + } + if systemType == "datetime" || systemType == "datetime2" || systemType == "timestamp" || systemType == "timestamptz" || systemType == "timestamp with time zone" { + return "TIMESTAMP" + } + + if systemType == "binary" || systemType == "varbinary" || systemType == "image" || systemType == "bytea" { + return "BINARY" + } + + if systemType == "uniqueidentifier" || systemType == "uuid" { + return "UUID" + } + + if systemType == "json" { + return "JSON" + } + + if systemType == "geometry" || systemType == "geography" { + return "GEOMETRY" + } + + return strings.ToUpper(systemType) +} + +func (ta *PostgresTableAnalyzer) rawColumnToColumnType(rawColumn rawColumnPostgres) models.ColumnType { + const nullValue int64 = -1 + stringTypes := map[string]bool{"varchar": true, "char": true, "text": true} + decimalTypes := map[string]bool{"decimal": true, "numeric": true} + + if stringTypes[rawColumn.systemType] { + rawColumn.precision, rawColumn.scale = nullValue, nullValue + } else if decimalTypes[rawColumn.systemType] { + rawColumn.maxLength = nullValue + } else { + rawColumn.maxLength, rawColumn.precision, rawColumn.scale = nullValue, nullValue, nullValue + } + + return models.NewColumnType( + rawColumn.name, + rawColumn.maxLength != nullValue, + rawColumn.precision != nullValue || rawColumn.scale != nullValue, + rawColumn.userType, + rawColumn.systemType, + ta.systemTypeToUnifiedType(rawColumn.systemType), + rawColumn.nullable, + rawColumn.maxLength, + rawColumn.precision, + rawColumn.scale, + ) +} + +func (ta *PostgresTableAnalyzer) QueryColumnTypes( + ctx context.Context, + tableInfo config.TableInfo, +) ([]models.ColumnType, error) { + localCtx, cancel := context.WithTimeout(ctx, 20*time.Second) + defer cancel() + + rows, err := ta.db.Query(localCtx, postgresColumnMetadataQuery, tableInfo.Schema, tableInfo.Table) + if err != nil { + return nil, err + } + defer rows.Close() + + var colTypes []models.ColumnType + + for rows.Next() { + var column rawColumnPostgres + + if err := rows.Scan( + &column.name, + &column.userType, + &column.systemType, + &column.nullable, + &column.maxLength, + &column.precision, + &column.scale, + ); err != nil { + return nil, err + } + + colTypes = append(colTypes, ta.rawColumnToColumnType(column)) + } + + return colTypes, nil +} + +func (ta *PostgresTableAnalyzer) EstimateTotalRows( + ctx context.Context, + tableInfo config.TableInfo, +) (int64, error) { + return 0, nil +} + +func (ta *PostgresTableAnalyzer) CalculatePartitionRanges( + ctx context.Context, + tableInfo config.TableInfo, + partitionColumn string, + maxPartitions int64, +) ([]models.Partition, error) { + return []models.Partition{}, nil +} diff --git a/internal/app/etl/transformers/mssql.go b/internal/app/etl/transformers/mssql.go index 23dd68e..7270ebb 100644 --- a/internal/app/etl/transformers/mssql.go +++ b/internal/app/etl/transformers/mssql.go @@ -60,15 +60,15 @@ func computeTransformationPlan(columns []models.ColumnType) []etl.ColumnTransfor return plan } -const processChunkCtxCheck = 4096 +const processBatchCtxCheck = 4096 -func (mssqlTr *MssqlTransformer) ProcessChunk( +func (mssqlTr *MssqlTransformer) ProcessBatch( ctx context.Context, - chunk *models.Batch, + batch *models.Batch, transformationPlan []etl.ColumnTransformPlan, ) error { - for i, rowValues := range chunk.Data { - if i%processChunkCtxCheck == 0 { + for i, rowValues := range batch.Rows { + if i%processBatchCtxCheck == 0 { if err := ctx.Err(); err != nil { return err } @@ -94,10 +94,10 @@ func (mssqlTr *MssqlTransformer) ProcessChunk( func (mssqlTr *MssqlTransformer) Exec( ctx context.Context, columns []models.ColumnType, - chChunksIn <-chan models.Batch, - chChunksOut chan<- models.Batch, + chBatchesIn <-chan models.Batch, + chBatchesOut chan<- models.Batch, chJobErrorsOut chan<- custom_errors.JobError, - wgActiveChunks *sync.WaitGroup, + wgActiveBatches *sync.WaitGroup, ) { transformationPlan := computeTransformationPlan(columns) @@ -110,22 +110,22 @@ func (mssqlTr *MssqlTransformer) Exec( case <-ctx.Done(): return - case chunk, ok := <-chChunksIn: + case batch, ok := <-chBatchesIn: if !ok { return } if len(transformationPlan) == 0 { select { - case chChunksOut <- chunk: - wgActiveChunks.Add(1) + case chBatchesOut <- batch: + wgActiveBatches.Add(1) continue case <-ctx.Done(): return } } - err := mssqlTr.ProcessChunk(ctx, &chunk, transformationPlan) + err := mssqlTr.ProcessBatch(ctx, &batch, transformationPlan) if err != nil { if errors.Is(err, ctx.Err()) { return @@ -139,12 +139,12 @@ func (mssqlTr *MssqlTransformer) Exec( } select { - case chChunksOut <- chunk: + case chBatchesOut <- batch: case <-ctx.Done(): return } - wgActiveChunks.Add(1) + wgActiveBatches.Add(1) } } } diff --git a/internal/app/etl/types.go b/internal/app/etl/types.go index 05fd79d..a2c07a5 100644 --- a/internal/app/etl/types.go +++ b/internal/app/etl/types.go @@ -10,14 +10,14 @@ import ( ) type Extractor interface { - ProcessBatch( + ProcessPartition( ctx context.Context, tableInfo config.SourceTableInfo, columns []models.ColumnType, - chunkSize int, - batch models.Partition, + batchSize int, + partition models.Partition, indexPrimaryKey int, - chChunksOut chan<- models.Batch, + chBatchesOut chan<- models.Batch, rowsRead *int64, ) error @@ -25,12 +25,12 @@ type Extractor interface { ctx context.Context, tableInfo config.SourceTableInfo, columns []models.ColumnType, - chunkSize int, - chBatchesIn <-chan models.Partition, - chChunksOut chan<- models.Batch, + batchSize int, + chPartitionsIn <-chan models.Partition, + chBatchesOut chan<- models.Batch, chErrorsOut chan<- custom_errors.ExtractorError, chJobErrorsOut chan<- custom_errors.JobError, - wgActiveBatches *sync.WaitGroup, + wgActivePartitions *sync.WaitGroup, rowsRead *int64, ) } @@ -43,43 +43,43 @@ type ColumnTransformPlan struct { } type Transformer interface { - ProcessChunk( + ProcessBatch( ctx context.Context, - chunk *models.Batch, + batch *models.Batch, transformationPlan []ColumnTransformPlan, ) error Exec( ctx context.Context, columns []models.ColumnType, - chChunksIn <-chan models.Batch, - chChunksOut chan<- models.Batch, + chBatchesIn <-chan models.Batch, + chBactchesOut chan<- models.Batch, chJobErrorsOut chan<- custom_errors.JobError, - wgActiveChunks *sync.WaitGroup, + wgActiveBatches *sync.WaitGroup, ) } type Loader interface { - ProcessChunk( + ProcessBatch( ctx context.Context, tableInfo config.TargetTableInfo, colNames []string, - chunk models.Batch, + batch models.Batch, ) (int, error) Exec( ctx context.Context, tableInfo config.TargetTableInfo, columns []models.ColumnType, - chChunksIn <-chan models.Batch, + chBatchesIn <-chan models.Batch, chErrorsOut chan<- custom_errors.LoaderError, chJobErrorsOut chan<- custom_errors.JobError, - wgActiveChunks *sync.WaitGroup, + wgActiveBatches *sync.WaitGroup, rowsLoaded *int64, ) } -type TableAnalizer interface { +type TableAnalyzer interface { QueryColumnTypes( ctx context.Context, tableInfo config.TableInfo, @@ -93,6 +93,7 @@ type TableAnalizer interface { CalculatePartitionRanges( ctx context.Context, tableInfo config.TableInfo, - totalPartitions int, - ) (models.Partition, error) + partitionColumn string, + maxPartitions int64, + ) ([]models.Partition, error) } diff --git a/internal/app/models/main.go b/internal/app/models/main.go index 6be796e..6eafb8d 100644 --- a/internal/app/models/main.go +++ b/internal/app/models/main.go @@ -7,7 +7,7 @@ type UnknownRowValues = []any type Batch struct { Id uuid.UUID PartitionId uuid.UUID - Data []UnknownRowValues + Rows []UnknownRowValues RetryCounter int }