diff --git a/examples/dynamoathenamigration/migration.go b/examples/dynamoathenamigration/migration.go index 952cc906720..3ed9a6a4878 100644 --- a/examples/dynamoathenamigration/migration.go +++ b/examples/dynamoathenamigration/migration.go @@ -38,6 +38,7 @@ import ( dynamoTypes "github.com/aws/aws-sdk-go-v2/service/dynamodb/types" "github.com/aws/aws-sdk-go-v2/service/s3" "github.com/aws/aws-sdk-go-v2/service/sns" + "github.com/google/uuid" "github.com/gravitational/trace" log "github.com/sirupsen/logrus" "golang.org/x/exp/maps" @@ -278,9 +279,11 @@ func (t *task) waitForCompletedExport(ctx context.Context, exportARN string) (ex select { case <-ctx.Done(): return "", trace.Wrap(ctx.Err()) - case <-time.After(10 * time.Second): + case <-time.After(30 * time.Second): t.Logger.Debug("Export job still in progress...") } + default: + return "", trace.Errorf("dynamo DescribeExport returned unexpected status: %v", exportStatus) } } @@ -410,6 +413,10 @@ func (t *task) fromS3ToChan(ctx context.Context, dataObj dataObjectInfo, eventsC return false, trace.Wrap(err) } + if ev.GetID() == "" || ev.GetID() == uuid.Nil.String() { + ev.SetID(uuid.NewString()) + } + // if checkpoint is present, it means that previous run ended with error // and we want to continue from last valid checkpoint. // We have list of checkpoints because processing is done in async way with @@ -433,7 +440,7 @@ func (t *task) fromS3ToChan(ctx context.Context, dataObj dataObjectInfo, eventsC return false, ctx.Err() } - if count%100 == 0 { + if count%1000 == 0 && !t.DryRun { t.Logger.Debugf("Sent on buffer %d/%d events from %s", count, dataObj.ItemCount, dataObj.DataFileS3Key) } } @@ -534,13 +541,22 @@ func (t *task) loadEmitterCheckpoint(ctx context.Context, exportARN string) (*ch return &out, nil } +type eventWithErr struct { + event apievents.AuditEvent + err error +} + func (t *task) emitEvents(ctx context.Context, eventsC <-chan apievents.AuditEvent, exportARN string) error { if t.DryRun { - // in dryRun we just want to count events, validation is done when reading from file. + var invalidEvents []eventWithErr var count int var oldest, newest apievents.AuditEvent for event := range eventsC { count++ + if validateErr := validateEvent(event); validateErr != nil { + invalidEvents = append(invalidEvents, eventWithErr{event: event, err: validateErr}) + continue + } if oldest == nil && newest == nil { // first iteration, initialize values with first event. oldest = event @@ -556,6 +572,12 @@ func (t *task) emitEvents(ctx context.Context, eventsC <-chan apievents.AuditEve if count == 0 { return errors.New("there were not events from export") } + if len(invalidEvents) > 0 { + for _, eventWithErr := range invalidEvents { + t.Logger.Debugf("Event %q %q %v is invalid: %v", eventWithErr.event.GetType(), eventWithErr.event.GetID(), eventWithErr.event.GetTime().Format(time.RFC3339), eventWithErr.err) + } + return trace.Errorf("there are %d invalid items", len(invalidEvents)) + } t.Logger.Infof("Dry run: there are %d events from %v to %v", count, oldest.GetTime(), newest.GetTime()) return nil } @@ -609,3 +631,18 @@ func (t *task) emitEvents(ctx context.Context, eventsC <-chan apievents.AuditEve } return trace.Wrap(workersErr) } + +func validateEvent(event apievents.AuditEvent) error { + if event.GetTime().IsZero() { + return trace.BadParameter("empty event time") + } + if _, err := uuid.Parse(event.GetID()); err != nil { + return trace.BadParameter("invalid uid format: %v", err) + } + oneOf, err := apievents.ToOneOf(event) + if err != nil { + return trace.Wrap(err) + } + _, err = oneOf.Marshal() + return trace.Wrap(err) +} diff --git a/examples/dynamoathenamigration/migration_test.go b/examples/dynamoathenamigration/migration_test.go index 3368d5d5df2..d94ec2bbb11 100644 --- a/examples/dynamoathenamigration/migration_test.go +++ b/examples/dynamoathenamigration/migration_test.go @@ -430,3 +430,83 @@ func generateDynamoExportData(n int) string { } return sb.String() } + +func TestMigrationDryRunValidation(t *testing.T) { + validEvent := func() apievents.AuditEvent { + return &apievents.AppCreate{ + Metadata: apievents.Metadata{ + Time: time.Date(2023, 5, 1, 12, 15, 0, 0, time.UTC), + ID: uuid.NewString(), + }, + } + } + tests := []struct { + name string + events func() []apievents.AuditEvent + wantLog string + wantErr string + }{ + { + name: "valid events", + events: func() []apievents.AuditEvent { + return []apievents.AuditEvent{ + validEvent(), validEvent(), + } + }, + }, + { + name: "event without time", + events: func() []apievents.AuditEvent { + eventWithoutTime := validEvent() + eventWithoutTime.SetTime(time.Time{}) + return []apievents.AuditEvent{ + validEvent(), eventWithoutTime, + } + }, + wantLog: "is invalid: empty event time", + wantErr: "1 invalid", + }, + { + name: "event with wrong uuid", + events: func() []apievents.AuditEvent { + eventWithInvalidUUID := validEvent() + eventWithInvalidUUID.SetID("invalid-uuid") + return []apievents.AuditEvent{ + validEvent(), eventWithInvalidUUID, + } + }, + wantLog: "is invalid: invalid uid format: invalid UUID length", + wantErr: "1 invalid", + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // Migration cli logs output from validation to logger. + var logBuffer bytes.Buffer + log := utils.NewLoggerForTests() + log.SetOutput(&logBuffer) + + tr := &task{ + Config: Config{ + Logger: log, + DryRun: true, + }, + } + c := make(chan apievents.AuditEvent, 10) + for _, e := range tt.events() { + c <- e + } + close(c) + err := tr.emitEvents(context.Background(), c, "" /* exportARN not used in dryRun */) + if tt.wantErr != "" { + require.ErrorContains(t, err, tt.wantErr) + } else { + require.NoError(t, err) + } + + if tt.wantLog != "" { + require.Contains(t, logBuffer.String(), tt.wantLog) + } + }) + } +}