package migration import ( "strings" "testing" ) func TestValidateForApplyAcceptsConfirmedSourceState(t *testing.T) { report := validMigrationReport() if err := ValidateForApply(report); err != nil { t.Fatalf("ValidateForApply() error = %v", err) } } func TestValidateForApplyRejectsUnsafeState(t *testing.T) { tests := []struct { name string mutate func(*Report) want string }{ { name: "read only connection", mutate: func(report *Report) { report.TransactionReadOnly = true }, want: "read-only", }, { name: "unexpected srid", mutate: func(report *Report) { report.Tables[0].UnexpectedSRIDCount = 1 }, want: "outside SRID", }, { name: "out of bounds", mutate: func(report *Report) { report.Tables[0].OutOfBoundsGeometryCount = 1 }, want: "outside WGS84 bounds", }, { name: "unexpected type", mutate: func(report *Report) { report.Tables[0].GeometryTypes = []string{"POLYGON"} }, want: "unexpected geometry type", }, { name: "cannot update", mutate: func(report *Report) { report.Tables[0].CanUpdate = false }, want: "cannot update", }, { name: "constrained source contains z", mutate: func(report *Report) { report.Tables[0].CoordinateDimensions = []int{2, 3} }, want: "cannot be safely retyped", }, { name: "cannot alter constrained column", mutate: func(report *Report) { report.Tables[0].CanAlter = false }, want: "cannot retype", }, } for _, test := range tests { t.Run(test.name, func(t *testing.T) { report := validMigrationReport() test.mutate(&report) err := ValidateForApply(report) if err == nil || !strings.Contains(err.Error(), test.want) { t.Fatalf("ValidateForApply() error = %v, want substring %q", err, test.want) } }) } } func TestValidateAfterAcceptsSRIDOnlyChange(t *testing.T) { before := validMigrationReport() after := migratedReport(before) if err := validateAfter(before, after); err != nil { t.Fatalf("validateAfter() error = %v", err) } } func TestValidateAfterRejectsGeometryPayloadChange(t *testing.T) { before := validMigrationReport() after := migratedReport(before) after.Tables[2].ZGeometryCount-- err := validateAfter(before, after) if err == nil || !strings.Contains(err.Error(), "geometry payload changed") { t.Fatalf("validateAfter() error = %v", err) } } func validMigrationReport() Report { report := Report{TargetSRID: TargetSRID} for index, spec := range sourceTables { state := TableState{ Table: spec.name, ColumnType: "geometry", TypmodConstrained: true, TypmodSRID: SourceSRID, TypmodDimensions: 2, TypmodGeometryType: "Geometry", RowCount: 3, NonNullGeometryCount: 2, NullGeometryCount: 1, SourceSRIDCount: 2, GeometryTypes: []string{spec.geometryTypes[0]}, SRIDs: []int{SourceSRID}, CoordinateDimensions: []int{2}, CanUpdate: true, CanAlter: true, fingerprint: "same-payload", } if index == 2 { state.TypmodConstrained = false state.TypmodSRID = -1 state.TypmodDimensions = -1 state.TypmodGeometryType = "" state.GeometryTypes = []string{"LINESTRING", "MULTILINESTRING"} state.CoordinateDimensions = []int{2, 3} state.ZGeometryCount = 1 } report.Tables = append(report.Tables, state) report.CandidateGeometryCount += state.SourceSRIDCount } return report } func migratedReport(before Report) Report { after := before after.Tables = append([]TableState(nil), before.Tables...) after.CandidateGeometryCount = 0 after.AlreadyTargetGeometryCount = 0 for index := range after.Tables { after.Tables[index].SourceSRIDCount = 0 after.Tables[index].TargetSRIDCount = after.Tables[index].NonNullGeometryCount after.Tables[index].SRIDs = []int{TargetSRID} after.AlreadyTargetGeometryCount += after.Tables[index].TargetSRIDCount } return after }