athenamigration: add validation on dry-run (#29193)

This commit is contained in:
Tobiasz Heller
2023-07-18 10:29:14 +00:00
committed by GitHub
parent c519c51378
commit 040ec6d3b2
2 changed files with 120 additions and 3 deletions
+40 -3
View File
@@ -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)
}
@@ -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)
}
})
}
}