package alertcontrol import ( "context" "encoding/json" "errors" "fmt" "time" "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgconn" "github.com/jackc/pgx/v5/pgxpool" ) type Store interface { CreateSilence(context.Context, string, Silence) (Silence, error) ListSilences(context.Context, int, time.Time) ([]Silence, error) RevokeSilence(context.Context, string, string, int64, time.Time) (Silence, error) CreateMaintenance(context.Context, string, MaintenanceWindow) (MaintenanceWindow, error) ListMaintenance(context.Context, int, time.Time) ([]MaintenanceWindow, error) RevokeMaintenance(context.Context, string, string, int64, time.Time) (MaintenanceWindow, error) Expire(context.Context, time.Time) (ExpiryResult, error) } type Repository struct{ Pool *pgxpool.Pool } type ExpiryResult struct { Silences int `json:"silences"` MaintenanceWindows int `json:"maintenanceWindows"` } func (r Repository) CreateSilence(ctx context.Context, actor string, silence Silence) (Silence, error) { if r.Pool == nil { return Silence{}, ErrUnavailable } if actor == "" { return Silence{}, fmt.Errorf("%w: creator is required", ErrInvalid) } if silence.ID == "" { silence.ID = NewID() } if silence.Owner == "" { silence.Owner = actor } if err := silence.Validate(time.Now().UTC()); err != nil { return Silence{}, err } matcher, err := json.Marshal(silence.Matchers) if err != nil { return Silence{}, fmt.Errorf("marshal silence matcher: %w", err) } if silence.Owner == "" { silence.Owner = actor } tx, err := r.Pool.BeginTx(ctx, pgx.TxOptions{}) if err != nil { return Silence{}, fmt.Errorf("begin silence create: %w", err) } defer func() { _ = tx.Rollback(ctx) }() _, err = tx.Exec(ctx, `INSERT INTO alert_silences (id,name,reason,owner,matchers,starts_at,expires_at,created_by) VALUES ($1::uuid,$2,$3,$4,$5::jsonb,$6,$7,$8)`, silence.ID, silence.Name, silence.Reason, silence.Owner, matcher, silence.StartsAt.UTC(), silence.ExpiresAt.UTC(), actor) if err != nil { return Silence{}, mapError(fmt.Errorf("create silence: %w", err)) } if err := tx.Commit(ctx); err != nil { return Silence{}, fmt.Errorf("commit silence create: %w", err) } return r.getSilence(ctx, silence.ID, time.Now().UTC()) } func (r Repository) ListSilences(ctx context.Context, limit int, now time.Time) ([]Silence, error) { if r.Pool == nil { return nil, ErrUnavailable } if limit < 1 || limit > 100 { return nil, fmt.Errorf("%w: invalid list limit", ErrInvalid) } rows, err := r.Pool.Query(ctx, `SELECT id,name,reason,owner,matchers,starts_at,expires_at,status,created_by,created_at,revoked_by,revoked_at,expired_at,revision FROM alert_silences ORDER BY starts_at DESC,id DESC LIMIT $1`, limit) if err != nil { return nil, fmt.Errorf("list silences: %w", err) } defer rows.Close() items := make([]Silence, 0) for rows.Next() { item, err := scanSilence(rows) if err != nil { return nil, err } item.State = item.StateAt(now.UTC()) items = append(items, item) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("iterate silences: %w", err) } return items, nil } func (r Repository) RevokeSilence(ctx context.Context, id, actor string, expected int64, now time.Time) (Silence, error) { if r.Pool == nil { return Silence{}, ErrUnavailable } if id == "" || actor == "" { return Silence{}, fmt.Errorf("%w: id and actor are required", ErrInvalid) } tx, err := r.Pool.BeginTx(ctx, pgx.TxOptions{}) if err != nil { return Silence{}, fmt.Errorf("begin silence revoke: %w", err) } defer func() { _ = tx.Rollback(ctx) }() var query string var args []any if expected > 0 { query = `UPDATE alert_silences SET status='revoked',revoked_by=$2,revoked_at=$3,revision=revision+1 WHERE id=$1 AND status='active' AND revision=$4` args = []any{id, actor, now.UTC(), expected} } else { query = `UPDATE alert_silences SET status='revoked',revoked_by=$2,revoked_at=$3,revision=revision+1 WHERE id=$1 AND status='active'` args = []any{id, actor, now.UTC()} } result, err := tx.Exec(ctx, query, args...) if err != nil { return Silence{}, mapError(fmt.Errorf("revoke silence: %w", err)) } if result.RowsAffected() == 0 { return Silence{}, r.revokeFailure(ctx, tx, id, expected, true) } if err := tx.Commit(ctx); err != nil { return Silence{}, fmt.Errorf("commit silence revoke: %w", err) } return r.getSilence(ctx, id, now.UTC()) } func (r Repository) CreateMaintenance(ctx context.Context, actor string, window MaintenanceWindow) (MaintenanceWindow, error) { if r.Pool == nil { return MaintenanceWindow{}, ErrUnavailable } if actor == "" { return MaintenanceWindow{}, fmt.Errorf("%w: creator is required", ErrInvalid) } if window.ID == "" { window.ID = NewID() } if err := window.Validate(time.Now().UTC()); err != nil { return MaintenanceWindow{}, err } selector, err := json.Marshal(window.Selector) if err != nil { return MaintenanceWindow{}, fmt.Errorf("marshal maintenance selector: %w", err) } tx, err := r.Pool.BeginTx(ctx, pgx.TxOptions{}) if err != nil { return MaintenanceWindow{}, fmt.Errorf("begin maintenance create: %w", err) } defer func() { _ = tx.Rollback(ctx) }() _, err = tx.Exec(ctx, `INSERT INTO maintenance_windows (id,name,reason,selector,starts_at,ends_at,created_by) VALUES ($1::uuid,$2,$3,$4::jsonb,$5,$6,$7)`, window.ID, window.Name, window.Reason, selector, window.StartsAt.UTC(), window.EndsAt.UTC(), actor) if err != nil { return MaintenanceWindow{}, mapError(fmt.Errorf("create maintenance window: %w", err)) } if err := tx.Commit(ctx); err != nil { return MaintenanceWindow{}, fmt.Errorf("commit maintenance create: %w", err) } return r.getMaintenance(ctx, window.ID, time.Now().UTC()) } func (r Repository) ListMaintenance(ctx context.Context, limit int, now time.Time) ([]MaintenanceWindow, error) { if r.Pool == nil { return nil, ErrUnavailable } if limit < 1 || limit > 100 { return nil, fmt.Errorf("%w: invalid list limit", ErrInvalid) } rows, err := r.Pool.Query(ctx, `SELECT id,name,reason,selector,starts_at,ends_at,status,created_by,created_at,revoked_by,revoked_at,expired_at,revision FROM maintenance_windows ORDER BY starts_at DESC,id DESC LIMIT $1`, limit) if err != nil { return nil, fmt.Errorf("list maintenance windows: %w", err) } defer rows.Close() items := make([]MaintenanceWindow, 0) for rows.Next() { item, err := scanMaintenance(rows) if err != nil { return nil, err } item.State = item.StateAt(now.UTC()) items = append(items, item) } if err := rows.Err(); err != nil { return nil, fmt.Errorf("iterate maintenance windows: %w", err) } return items, nil } func (r Repository) RevokeMaintenance(ctx context.Context, id, actor string, expected int64, now time.Time) (MaintenanceWindow, error) { if r.Pool == nil { return MaintenanceWindow{}, ErrUnavailable } if id == "" || actor == "" { return MaintenanceWindow{}, fmt.Errorf("%w: id and actor are required", ErrInvalid) } tx, err := r.Pool.BeginTx(ctx, pgx.TxOptions{}) if err != nil { return MaintenanceWindow{}, fmt.Errorf("begin maintenance revoke: %w", err) } defer func() { _ = tx.Rollback(ctx) }() var query string var args []any if expected > 0 { query = `UPDATE maintenance_windows SET status='revoked',revoked_by=$2,revoked_at=$3,revision=revision+1 WHERE id=$1 AND status='active' AND revision=$4` args = []any{id, actor, now.UTC(), expected} } else { query = `UPDATE maintenance_windows SET status='revoked',revoked_by=$2,revoked_at=$3,revision=revision+1 WHERE id=$1 AND status='active'` args = []any{id, actor, now.UTC()} } result, err := tx.Exec(ctx, query, args...) if err != nil { return MaintenanceWindow{}, mapError(fmt.Errorf("revoke maintenance window: %w", err)) } if result.RowsAffected() == 0 { return MaintenanceWindow{}, r.revokeFailure(ctx, tx, id, expected, false) } if err := tx.Commit(ctx); err != nil { return MaintenanceWindow{}, fmt.Errorf("commit maintenance revoke: %w", err) } return r.getMaintenance(ctx, id, now.UTC()) } func (r Repository) Expire(ctx context.Context, now time.Time) (ExpiryResult, error) { if r.Pool == nil { return ExpiryResult{}, ErrUnavailable } now = now.UTC() tx, err := r.Pool.BeginTx(ctx, pgx.TxOptions{}) if err != nil { return ExpiryResult{}, fmt.Errorf("begin control expiry: %w", err) } defer func() { _ = tx.Rollback(ctx) }() var result ExpiryResult if err := tx.QueryRow(ctx, `WITH expired AS (UPDATE alert_silences SET status='expired',expired_at=$1,revision=revision+1 WHERE status='active' AND expires_at <= $1 RETURNING id) SELECT count(*) FROM expired`, now).Scan(&result.Silences); err != nil { return ExpiryResult{}, fmt.Errorf("expire silences: %w", err) } if err := tx.QueryRow(ctx, `WITH expired AS (UPDATE maintenance_windows SET status='expired',expired_at=$1,revision=revision+1 WHERE status='active' AND ends_at <= $1 RETURNING id) SELECT count(*) FROM expired`, now).Scan(&result.MaintenanceWindows); err != nil { return ExpiryResult{}, fmt.Errorf("expire maintenance windows: %w", err) } if err := tx.Commit(ctx); err != nil { return ExpiryResult{}, fmt.Errorf("commit control expiry: %w", err) } return result, nil } func (r Repository) getSilence(ctx context.Context, id string, now time.Time) (Silence, error) { var item Silence row := r.Pool.QueryRow(ctx, `SELECT id,name,reason,owner,matchers,starts_at,expires_at,status,created_by,created_at,revoked_by,revoked_at,expired_at,revision FROM alert_silences WHERE id=$1`, id) scanned, err := scanSilence(row) if errors.Is(err, pgx.ErrNoRows) { return Silence{}, ErrNotFound } if err != nil { return Silence{}, fmt.Errorf("get silence: %w", err) } item = scanned item.State = item.StateAt(now.UTC()) return item, nil } func (r Repository) getMaintenance(ctx context.Context, id string, now time.Time) (MaintenanceWindow, error) { row := r.Pool.QueryRow(ctx, `SELECT id,name,reason,selector,starts_at,ends_at,status,created_by,created_at,revoked_by,revoked_at,expired_at,revision FROM maintenance_windows WHERE id=$1`, id) item, err := scanMaintenance(row) if errors.Is(err, pgx.ErrNoRows) { return MaintenanceWindow{}, ErrNotFound } if err != nil { return MaintenanceWindow{}, fmt.Errorf("get maintenance window: %w", err) } item.State = item.StateAt(now.UTC()) return item, nil } type rowScanner interface{ Scan(...any) error } func scanSilence(row rowScanner) (Silence, error) { var item Silence var matcher []byte var status string var revokedBy, createdBy *string if err := row.Scan(&item.ID, &item.Name, &item.Reason, &item.Owner, &matcher, &item.StartsAt, &item.ExpiresAt, &status, &createdBy, &item.CreatedAt, &revokedBy, &item.RevokedAt, &item.ExpiredAt, &item.Revision); err != nil { return Silence{}, err } if err := json.Unmarshal(matcher, &item.Matchers); err != nil { return Silence{}, fmt.Errorf("decode silence matcher: %w", err) } item.CreatedBy = valueOrEmpty(createdBy) item.RevokedBy = valueOrEmpty(revokedBy) return item, nil } func scanMaintenance(row rowScanner) (MaintenanceWindow, error) { var item MaintenanceWindow var selector []byte var status string var revokedBy, createdBy *string if err := row.Scan(&item.ID, &item.Name, &item.Reason, &selector, &item.StartsAt, &item.EndsAt, &status, &createdBy, &item.CreatedAt, &revokedBy, &item.RevokedAt, &item.ExpiredAt, &item.Revision); err != nil { return MaintenanceWindow{}, err } if err := json.Unmarshal(selector, &item.Selector); err != nil { return MaintenanceWindow{}, fmt.Errorf("decode maintenance selector: %w", err) } item.CreatedBy = valueOrEmpty(createdBy) item.RevokedBy = valueOrEmpty(revokedBy) return item, nil } func valueOrEmpty(value *string) string { if value == nil { return "" } return *value } func (r Repository) revokeFailure(ctx context.Context, tx pgx.Tx, id string, expected int64, silence bool) error { var revision int64 var status string table := "maintenance_windows" if silence { table = "alert_silences" } err := tx.QueryRow(ctx, "SELECT revision,status FROM "+table+" WHERE id=$1 FOR UPDATE", id).Scan(&revision, &status) if errors.Is(err, pgx.ErrNoRows) { return ErrNotFound } if err != nil { return fmt.Errorf("inspect control revoke: %w", err) } if expected > 0 && revision != expected { return ErrConflict } return fmt.Errorf("%w: control is %s", ErrConflict, status) } func mapError(err error) error { var pgErr *pgconn.PgError if errors.As(err, &pgErr) { switch pgErr.Code { case "23505": return fmt.Errorf("%w: duplicate control", ErrConflict) case "23514", "22P02": return fmt.Errorf("%w: database constraint", ErrInvalid) } } return err }