package audit import ( "context" "errors" "fmt" "regexp" "sort" "time" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" ) const maintenanceLockID int64 = 6720240816 var auditPartitionPattern = regexp.MustCompile(`^audit_events_(\d{4})(\d{2})$`) type MaintenanceResult struct { CreatedPartitions []string `json:"created_partitions"` DroppedPartitions []string `json:"dropped_partitions"` DeletedAuditRows int64 `json:"deleted_audit_rows"` DeletedUsageRows int64 `json:"deleted_usage_rows"` } type Maintenance struct { pool *pgxpool.Pool auditRetention time.Duration usageRetention time.Duration monthsAhead int } func NewMaintenance(pool *pgxpool.Pool, auditRetention, usageRetention time.Duration, monthsAhead int) *Maintenance { return &Maintenance{pool: pool, auditRetention: auditRetention, usageRetention: usageRetention, monthsAhead: monthsAhead} } func (m *Maintenance) Run(ctx context.Context, now time.Time) (MaintenanceResult, error) { var result MaintenanceResult if m == nil || m.pool == nil { return result, errors.New("audit maintenance store unavailable") } now = now.UTC() auditCutoff := now.Add(-m.auditRetention) usageCutoff := now.Add(-m.usageRetention) tx, err := m.pool.Begin(ctx) if err != nil { return result, fmt.Errorf("begin audit maintenance: %w", err) } defer func() { _ = tx.Rollback(ctx) }() if _, err := tx.Exec(ctx, `SELECT pg_advisory_xact_lock($1)`, maintenanceLockID); err != nil { return result, fmt.Errorf("lock audit maintenance: %w", err) } if _, err := tx.Exec(ctx, `LOCK TABLE gateway.audit_events IN ACCESS EXCLUSIVE MODE`); err != nil { return result, fmt.Errorf("lock audit table: %w", err) } partitions, err := listAuditPartitions(ctx, tx) if err != nil { return result, err } for name := range partitions { start, ok := auditPartitionMonth(name) if !ok || start.AddDate(0, 1, 0).After(auditCutoff) { continue } if _, err := tx.Exec(ctx, `DROP TABLE `+pgx.Identifier{"gateway", name}.Sanitize()); err != nil { return result, fmt.Errorf("drop audit partition %s: %w", name, err) } result.DroppedPartitions = append(result.DroppedPartitions, name) } if _, err := tx.Exec(ctx, `ALTER TABLE gateway.audit_events DETACH PARTITION gateway.audit_events_default`); err != nil { return result, fmt.Errorf("detach audit default partition: %w", err) } deleted, err := tx.Exec(ctx, `DELETE FROM gateway.audit_events_default WHERE recorded_at < $1`, auditCutoff) if err != nil { return result, fmt.Errorf("clean audit default partition: %w", err) } result.DeletedAuditRows += deleted.RowsAffected() from := monthStart(auditCutoff) through := monthStart(now).AddDate(0, m.monthsAhead+1, 0) for start := from; start.Before(through); start = start.AddDate(0, 1, 0) { end := start.AddDate(0, 1, 0) name := "audit_events_" + start.Format("200601") if _, exists := partitions[name]; !exists { statement := fmt.Sprintf(`CREATE TABLE %s PARTITION OF gateway.audit_events FOR VALUES FROM ('%s') TO ('%s')`, pgx.Identifier{"gateway", name}.Sanitize(), start.Format(time.RFC3339), end.Format(time.RFC3339)) if _, err := tx.Exec(ctx, statement); err != nil { return result, fmt.Errorf("create audit partition %s: %w", name, err) } result.CreatedPartitions = append(result.CreatedPartitions, name) } statement := `WITH moved AS (DELETE FROM gateway.audit_events_default WHERE recorded_at >= $1 AND recorded_at < $2 RETURNING *) INSERT INTO gateway.audit_events SELECT * FROM moved` if _, err := tx.Exec(ctx, statement, start, end); err != nil { return result, fmt.Errorf("move default audit rows into %s: %w", name, err) } } if _, err := tx.Exec(ctx, `ALTER TABLE gateway.audit_events ATTACH PARTITION gateway.audit_events_default DEFAULT`); err != nil { return result, fmt.Errorf("reattach audit default partition: %w", err) } deleted, err = tx.Exec(ctx, `DELETE FROM gateway.audit_events WHERE recorded_at < $1`, auditCutoff) if err != nil { return result, fmt.Errorf("apply exact audit retention: %w", err) } result.DeletedAuditRows += deleted.RowsAffected() deleted, err = tx.Exec(ctx, `DELETE FROM gateway.usage_daily WHERE usage_date < $1::date`, usageCutoff.Format("2006-01-02")) if err != nil { return result, fmt.Errorf("apply usage retention: %w", err) } result.DeletedUsageRows = deleted.RowsAffected() if err := tx.Commit(ctx); err != nil { return result, fmt.Errorf("commit audit maintenance: %w", err) } sort.Strings(result.CreatedPartitions) sort.Strings(result.DroppedPartitions) return result, nil } func listAuditPartitions(ctx context.Context, tx pgx.Tx) (map[string]struct{}, error) { rows, err := tx.Query(ctx, `SELECT child.relname FROM pg_inherits i JOIN pg_class parent ON parent.oid=i.inhparent JOIN pg_namespace n ON n.oid=parent.relnamespace JOIN pg_class child ON child.oid=i.inhrelid WHERE n.nspname='gateway' AND parent.relname='audit_events'`) if err != nil { return nil, fmt.Errorf("list audit partitions: %w", err) } defer rows.Close() result := make(map[string]struct{}) for rows.Next() { var name string if err := rows.Scan(&name); err != nil { return nil, fmt.Errorf("scan audit partition: %w", err) } if name != "audit_events_default" { result[name] = struct{}{} } } return result, rows.Err() } func auditPartitionMonth(name string) (time.Time, bool) { match := auditPartitionPattern.FindStringSubmatch(name) if match == nil { return time.Time{}, false } parsed, err := time.Parse("200601", match[1]+match[2]) return parsed.UTC(), err == nil } func monthStart(value time.Time) time.Time { value = value.UTC() return time.Date(value.Year(), value.Month(), 1, 0, 0, 0, 0, time.UTC) }