package database_test import ( "context" "net/http" "testing" "time" "github.com/google/uuid" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "gorm.io/gorm" "sneak.berlin/go/webhooker/internal/database" ) // readEventTotals reads a webhook database's row of event totals, // asserting that it has exactly one. func readEventTotals(t *testing.T, db *gorm.DB) database.EventTotals { t.Helper() var rows []database.EventTotals require.NoError(t, db.Find(&rows).Error) require.Len(t, rows, 1) return rows[0] } // readTargetTotals reads a webhook database's target totals, keyed by // target. func readTargetTotals( t *testing.T, db *gorm.DB, ) map[string]database.TargetTotals { t.Helper() var rows []database.TargetTotals require.NoError(t, db.Find(&rows).Error) byTarget := make(map[string]database.TargetTotals, len(rows)) for _, row := range rows { byTarget[row.TargetID] = row } return byTarget } // TestWebhookDBManager_TotalsSurviveReopen verifies that a new event // database starts with one row of zero event totals and no target // totals, that adding to a target twice adds to the one row, and that // opening the database again keeps everything added. func TestWebhookDBManager_TotalsSurviveReopen(t *testing.T) { t.Parallel() mgr, lc := setupTestWebhookDBManager(t) ctx := context.Background() require.NoError(t, lc.Start(ctx)) defer func() { require.NoError(t, lc.Stop(ctx)) }() webhookID := uuid.New().String() db, err := mgr.GetDB(webhookID) require.NoError(t, err) fresh := readEventTotals(t, db) assert.Equal(t, database.EventTotals{ID: fresh.ID}, fresh) assert.Empty(t, readTargetTotals(t, db)) first, second := uuid.New().String(), uuid.New().String() require.NoError(t, database.AddEventTotals(db, database.EventTotals{ Events: 2, })) require.NoError(t, database.AddTargetTotals(db, database.TargetTotals{ TargetID: first, Deliveries: 2, Delivered: 1, })) require.NoError(t, database.AddTargetTotals(db, database.TargetTotals{ TargetID: first, Failed: 1, })) require.NoError(t, database.AddTargetTotals(db, database.TargetTotals{ TargetID: second, Deliveries: 1, })) // Drop the cached connection so the next open reopens the file, // as a restart would. require.NoError(t, mgr.CloseAll()) db, err = mgr.GetDB(webhookID) require.NoError(t, err) assert.Equal(t, database.EventTotals{ID: fresh.ID, Events: 2}, readEventTotals(t, db)) assert.Equal(t, map[string]database.TargetTotals{ first: { TargetID: first, Deliveries: 2, Delivered: 1, Failed: 1, }, second: {TargetID: second, Deliveries: 1}, }, readTargetTotals(t, db)) } // seedExpiredEvents stores count events created at the given time, // each with a delivered delivery to one target and a failed delivery // to the other, and one attempt for each delivery. func seedExpiredEvents( t *testing.T, db *gorm.DB, webhookID string, count int, createdAt time.Time, delivered, failed string, ) { t.Helper() events := make([]database.Event, count) deliveries := make([]database.Delivery, 0, 2*count) for i := range events { events[i] = database.Event{ WebhookID: webhookID, EntrypointID: uuid.New().String(), Method: http.MethodPost, } events[i].ID = uuid.New().String() events[i].CreatedAt = createdAt deliveries = append(deliveries, database.Delivery{ EventID: events[i].ID, TargetID: delivered, Status: database.DeliveryStatusDelivered, }, database.Delivery{ EventID: events[i].ID, TargetID: failed, Status: database.DeliveryStatusFailed, }, ) } require.NoError(t, db.CreateInBatches(events, 500).Error) require.NoError(t, db.CreateInBatches(deliveries, 500).Error) results := make([]database.DeliveryResult, len(deliveries)) for i := range deliveries { results[i] = database.DeliveryResult{ DeliveryID: deliveries[i].ID, AttemptNum: 1, } } require.NoError(t, db.CreateInBatches(results, 500).Error) } // seedBareEvents stores count events created at the given time, with // no deliveries. func seedBareEvents( t *testing.T, db *gorm.DB, webhookID string, count int, createdAt time.Time, ) { t.Helper() events := make([]database.Event, count) for i := range events { events[i] = database.Event{ WebhookID: webhookID, EntrypointID: uuid.New().String(), Method: http.MethodPost, } events[i].CreatedAt = createdAt } require.NoError(t, db.CreateInBatches(events, 500).Error) } // TestRetentionReaper_PrunesMoreThanOneBatch verifies that a prune // larger than one transaction's batch removes every expired event with // its deliveries and delivery results, keeps the recent event, and // adds what it removed to the event and target totals, so the totals // within retention match the rows still stored. func TestRetentionReaper_PrunesMoreThanOneBatch(t *testing.T) { t.Parallel() env := setupRetentionTest(t) webhookID := createWebhook(t, env.mainDB.DB(), 30) db, err := env.mgr.GetDB(webhookID) require.NoError(t, err) expired := database.ExportReapBatchSize + 1 delivered, failed := uuid.New().String(), uuid.New().String() seedExpiredEvents(t, db, webhookID, expired, time.Now().Add(-40*24*time.Hour), delivered, failed) // One recent event, delivered to the first target. recent := seedEventChain(t, db, webhookID, time.Now()) require.NoError(t, db.Model(&database.Delivery{}). Where("id = ?", recent.deliveryID). Update("target_id", delivered).Error) // The totals storing those rows would have left. n := int64(expired) require.NoError(t, database.AddEventTotals(db, database.EventTotals{ Events: n + 1, })) require.NoError(t, database.AddTargetTotals(db, database.TargetTotals{ TargetID: delivered, Deliveries: n + 1, Delivered: n + 1, })) require.NoError(t, database.AddTargetTotals(db, database.TargetTotals{ TargetID: failed, Deliveries: n, Failed: n, })) env.reaper.ExportSweep(context.Background()) // Only the recent event's rows are left. for _, model := range []any{ &database.Event{}, &database.Delivery{}, &database.DeliveryResult{}, } { var count int64 require.NoError(t, db.Model(model).Count(&count).Error) assert.Equal(t, int64(1), count, "%T rows left", model) } assertChainPresent(t, db, recent) eventTotals := readEventTotals(t, db) assert.Equal(t, database.EventTotals{ ID: eventTotals.ID, Events: n + 1, EventsRemoved: n, }, eventTotals) targetTotals := readTargetTotals(t, db) assert.Equal(t, map[string]database.TargetTotals{ delivered: { TargetID: delivered, Deliveries: n + 1, Delivered: n + 1, DeliveriesRemoved: n, }, failed: { TargetID: failed, Deliveries: n, Failed: n, DeliveriesRemoved: n, FailedRemoved: n, }, }, targetTotals) // A sweep with nothing left to remove changes nothing. env.reaper.ExportSweep(context.Background()) assert.Equal(t, eventTotals, readEventTotals(t, db)) assert.Equal(t, targetTotals, readTargetTotals(t, db)) } // TestRetentionReaper_WriteDuringPruneSucceeds verifies that a prune // of several batches lets other writers in between its batches: an // event stored once the first batch is deleted is stored while expired // events are still left, not only after the prune has finished. func TestRetentionReaper_WriteDuringPruneSucceeds(t *testing.T) { t.Parallel() env := setupRetentionTest(t) webhookID := createWebhook(t, env.mainDB.DB(), 30) db, err := env.mgr.GetDB(webhookID) require.NoError(t, err) // Three batches of expired events, with nothing else stored: only // the number of batches matters here. expired := 3 * database.ExportReapBatchSize seedBareEvents(t, db, webhookID, expired, time.Now().Add(-40*24*time.Hour)) cutoff := time.Now().Add(-30 * 24 * time.Hour) countExpired := func() int64 { var count int64 require.NoError(t, db.Model(&database.Event{}). Where("created_at < ?", cutoff). Count(&count).Error) return count } pruned := make(chan struct{}) go func() { defer close(pruned) env.reaper.ExportSweep(context.Background()) }() t.Cleanup(func() { <-pruned }) // Every stored event is expired until the write below. require.Eventually(t, func() bool { var count int64 err := db.Model(&database.Event{}).Count(&count).Error return err == nil && count < int64(expired) }, 10*time.Second, 10*time.Millisecond) event := &database.Event{ WebhookID: webhookID, EntrypointID: uuid.New().String(), Method: http.MethodPost, } require.NoError(t, db.Create(event).Error) assert.Positive(t, countExpired(), "the event was stored only after the whole prune") <-pruned assert.Zero(t, countExpired()) var stored database.Event require.NoError(t, db.First(&stored, "id = ?", event.ID).Error) } // TestRetentionReaper_StopDuringPruneLeavesTheRest verifies that // stopping the reaper during a prune of several batches returns // between two batches, well inside the stop timeout, leaving the // remaining expired events for the next sweep, and that the totals // match the rows left. func TestRetentionReaper_StopDuringPruneLeavesTheRest(t *testing.T) { t.Parallel() env := setupRetentionTest(t) webhookID := createWebhook(t, env.mainDB.DB(), 30) db, err := env.mgr.GetDB(webhookID) require.NoError(t, err) // Two batches and one more of expired events, a few of them with a // delivered and a failed delivery for the target totals to count. // Most carry nothing else, to keep the test quick. const withDeliveries = 10 expiredAt := time.Now().Add(-40 * 24 * time.Hour) delivered, failed := uuid.New().String(), uuid.New().String() seedExpiredEvents(t, db, webhookID, withDeliveries, expiredAt, delivered, failed) seedBareEvents(t, db, webhookID, 2*database.ExportReapBatchSize+1-withDeliveries, expiredAt) n := int64(2*database.ExportReapBatchSize + 1) require.NoError(t, database.AddEventTotals(db, database.EventTotals{ Events: n, })) require.NoError(t, database.AddTargetTotals(db, database.TargetTotals{ TargetID: delivered, Deliveries: withDeliveries, Delivered: withDeliveries, })) require.NoError(t, database.AddTargetTotals(db, database.TargetTotals{ TargetID: failed, Deliveries: withDeliveries, Failed: withDeliveries, })) env.reaper.ExportSetInterval(time.Millisecond) env.reaper.ExportStart() // Stop once the first batch is deleted. The stop lands in the pause // after it, or at worst during the second batch, so at least the // last event is left. require.Eventually(t, func() bool { var count int64 err := db.Model(&database.Event{}).Count(&count).Error return err == nil && count < n }, 10*time.Second, 10*time.Millisecond) // The app's stop timeout. ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() require.NoError(t, env.reaper.ExportStop(ctx)) var events int64 require.NoError(t, db.Model(&database.Event{}).Count(&events).Error) assert.Positive(t, events, "the stop waited for the whole prune") eventTotals := readEventTotals(t, db) assert.Equal(t, events, eventTotals.Events-eventTotals.EventsRemoved) targetTotals := readTargetTotals(t, db) require.Len(t, targetTotals, 2) for target, totals := range targetTotals { var deliveries, failures int64 require.NoError(t, db.Model(&database.Delivery{}). Where("target_id = ?", target). Count(&deliveries).Error) require.NoError(t, db.Model(&database.Delivery{}). Where("target_id = ? AND status = ?", target, database.DeliveryStatusFailed). Count(&failures).Error) assert.Equal(t, deliveries, totals.Deliveries-totals.DeliveriesRemoved, target) assert.Equal(t, failures, totals.Failed-totals.FailedRemoved, target) } }