Files
cloud-ip-validator/internal/db/queries_runs.go
T

387 lines
12 KiB
Go
Raw Normal View History

package db
import (
"context"
"database/sql"
"fmt"
"net/netip"
"sort"
"time"
)
// Run kinds and states. A run groups the cycles of one launch of checks — see
// migrations/0011_check_runs.sql.
const (
RunManual = "manual"
RunAuto = "auto"
RunOpen = "open"
RunFinalized = "finalized"
)
// CheckRun is one launch of checks. FinalizedAt is nil while it is open.
type CheckRun struct {
ID int64
Kind string
State string
StartedAt time.Time
FinalizedAt *time.Time
}
// RunCounts is the verdict tally of a run.
type RunCounts struct {
Addresses int
Pass int
Partial int
Fail int
Cancelled int
}
// RunSummary is a run with its verdict tally, for the run selector. Pending
// is the number of its queue rows still being processed (open runs only).
type RunSummary struct {
CheckRun
RunCounts
Pending int
Total int
}
// openRunTx returns the id of the open run, creating one of the given kind if
// none is open. Called when addresses enter the queue: while a run is open
// everything submitted or re-checked joins it.
func openRunTx(ctx context.Context, tx *sql.Tx, kind, now string) (int64, error) {
var id int64
err := tx.QueryRowContext(ctx, `SELECT id FROM check_runs WHERE state=? ORDER BY id DESC LIMIT 1`, RunOpen).Scan(&id)
if err == nil {
return id, nil
}
if err != sql.ErrNoRows {
return 0, err
}
if kind == "" {
kind = RunManual
}
res, err := tx.ExecContext(ctx, `INSERT INTO check_runs (kind, state, started_at) VALUES (?, ?, ?)`, kind, RunOpen, now)
if err != nil {
return 0, fmt.Errorf("open run: %w", err)
}
return res.LastInsertId()
}
// finalizeRunsTx finalizes every open run that has no queue row left in a
// non-terminal state, and drops finalized runs that ended up with no results
// at all (everything was deleted before any verdict). Called wherever a queue
// row reaches a terminal state or disappears.
func finalizeRunsTx(ctx context.Context, tx *sql.Tx, now string) error {
if _, err := tx.ExecContext(ctx, `
UPDATE check_runs SET state=?, finalized_at=COALESCE(
(SELECT MAX(aggregated_at) FROM run_results WHERE run_id=check_runs.id), ?)
WHERE state=? AND NOT EXISTS (
SELECT 1 FROM ip_queue q WHERE q.run_id=check_runs.id AND q.state NOT IN (?, ?, ?))
`, RunFinalized, now, RunOpen, IPDone, IPFailed, IPOccupied); err != nil {
return fmt.Errorf("finalize runs: %w", err)
}
if _, err := tx.ExecContext(ctx, `
DELETE FROM check_runs WHERE state=? AND NOT EXISTS (SELECT 1 FROM run_results r WHERE r.run_id=check_runs.id)
`, RunFinalized); err != nil {
return fmt.Errorf("drop empty runs: %w", err)
}
return nil
}
// FinalizeRuns finalizes runs whose addresses are all done. The orchestrator
// calls it every tick as a safety net; the terminal transitions already do it.
func (d *DB) FinalizeRuns(ctx context.Context) error {
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if err := finalizeRunsTx(ctx, tx, timeToDB(Now())); err != nil {
return err
}
return tx.Commit()
}
// upsertRunResultTx stores the verdict of the address's current cycle in its
// run. A re-check inside an open run replaces the earlier result. expected < 0
// means unknown. Rows without a run (queue rows that predate runs and never
// got one) are skipped.
func upsertRunResultTx(ctx context.Context, tx *sql.Tx, ipID int64, verdict string, expected int, now string) error {
var exp sql.NullInt64
if expected >= 0 {
exp = sql.NullInt64{Int64: int64(expected), Valid: true}
}
_, err := tx.ExecContext(ctx, `
INSERT INTO run_results (run_id, registry_id, ip_address, cycle_id, verdict, aggregated_at, expected_checks, recorded_checks)
SELECT q.run_id, q.registry_id, q.ip_address, q.cycle_id, ?, ?, ?,
(SELECT COUNT(*) FROM checks c WHERE c.registry_id=q.registry_id AND c.cycle_id=q.cycle_id)
FROM ip_queue q WHERE q.id=? AND q.run_id IS NOT NULL
ON CONFLICT(run_id, registry_id) DO UPDATE SET
cycle_id=excluded.cycle_id, verdict=excluded.verdict, verdict_derived=0,
aggregated_at=excluded.aggregated_at, expected_checks=excluded.expected_checks,
recorded_checks=excluded.recorded_checks
`, verdict, now, exp, ipID)
if err != nil {
return fmt.Errorf("record run result: %w", err)
}
return nil
}
// adoptOrphanQueueRows gives every live queue row without a run (rows that
// predate runs) the open run, creating one. Runs once after migrating.
func (d *DB) adoptOrphanQueueRows(ctx context.Context) error {
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
var n int
if err := tx.QueryRowContext(ctx, `SELECT COUNT(*) FROM ip_queue WHERE run_id IS NULL AND state NOT IN (?, ?, ?)`,
IPDone, IPFailed, IPOccupied).Scan(&n); err != nil {
return err
}
if n == 0 {
return nil
}
now := timeToDB(Now())
id, err := openRunTx(ctx, tx, RunManual, now)
if err != nil {
return err
}
if _, err := tx.ExecContext(ctx, `UPDATE ip_queue SET run_id=? WHERE run_id IS NULL AND state NOT IN (?, ?, ?)`,
id, IPDone, IPFailed, IPOccupied); err != nil {
return err
}
return tx.Commit()
}
// ListRuns returns every run, newest first, with its verdict tally.
func (d *DB) ListRuns(ctx context.Context) ([]RunSummary, error) {
rows, err := d.QueryContext(ctx, `
SELECT r.id, r.kind, r.state, r.started_at, r.finalized_at,
COALESCE(SUM(CASE WHEN x.verdict IS NOT NULL THEN 1 ELSE 0 END), 0),
COALESCE(SUM(CASE WHEN x.verdict='pass' THEN 1 ELSE 0 END), 0),
COALESCE(SUM(CASE WHEN x.verdict='partial' THEN 1 ELSE 0 END), 0),
COALESCE(SUM(CASE WHEN x.verdict='fail' THEN 1 ELSE 0 END), 0),
COALESCE(SUM(CASE WHEN x.verdict='cancelled' THEN 1 ELSE 0 END), 0),
(SELECT COUNT(*) FROM ip_queue q WHERE q.run_id=r.id),
(SELECT COUNT(*) FROM ip_queue q WHERE q.run_id=r.id AND q.state NOT IN (?, ?, ?))
FROM check_runs r LEFT JOIN run_results x ON x.run_id=r.id
GROUP BY r.id ORDER BY r.id DESC
`, IPDone, IPFailed, IPOccupied)
if err != nil {
return nil, err
}
defer rows.Close()
var out []RunSummary
for rows.Next() {
var s RunSummary
var started string
var finalized sql.NullString
if err := rows.Scan(&s.ID, &s.Kind, &s.State, &started, &finalized,
&s.Addresses, &s.Pass, &s.Partial, &s.Fail, &s.Cancelled, &s.Total, &s.Pending); err != nil {
return nil, err
}
if s.StartedAt, err = dbToTime(started); err != nil {
return nil, err
}
if s.FinalizedAt, err = nullStringToTimePtr(finalized); err != nil {
return nil, err
}
out = append(out, s)
}
return out, rows.Err()
}
// GetRun returns one run, or ErrNotFound.
func (d *DB) GetRun(ctx context.Context, id int64) (*CheckRun, error) {
var r CheckRun
var started string
var finalized sql.NullString
err := d.QueryRowContext(ctx, `SELECT id, kind, state, started_at, finalized_at FROM check_runs WHERE id=?`, id).
Scan(&r.ID, &r.Kind, &r.State, &started, &finalized)
if err == sql.ErrNoRows {
return nil, fmt.Errorf("run %d: %w", id, ErrNotFound)
}
if err != nil {
return nil, err
}
if r.StartedAt, err = dbToTime(started); err != nil {
return nil, err
}
if r.FinalizedAt, err = nullStringToTimePtr(finalized); err != nil {
return nil, err
}
return &r, nil
}
// RunResult is one address's result in a run.
type RunResult struct {
RegistryID int64
IPAddress string
CycleID int
Verdict string
Derived bool
AggregatedAt time.Time
ExpectedChecks int // -1 when unknown
RecordedChecks int
}
// ListRunResults returns the results of a run in address order of insertion.
func (d *DB) ListRunResults(ctx context.Context, runID int64) ([]RunResult, error) {
rows, err := d.QueryContext(ctx, `
SELECT registry_id, ip_address, cycle_id, verdict, verdict_derived, aggregated_at, expected_checks, recorded_checks
FROM run_results WHERE run_id=? ORDER BY registry_id`, runID)
if err != nil {
return nil, err
}
defer rows.Close()
var out []RunResult
for rows.Next() {
var r RunResult
var agg string
var exp sql.NullInt64
if err := rows.Scan(&r.RegistryID, &r.IPAddress, &r.CycleID, &r.Verdict, &r.Derived, &agg, &exp, &r.RecordedChecks); err != nil {
return nil, err
}
if r.AggregatedAt, err = dbToTime(agg); err != nil {
return nil, err
}
r.ExpectedChecks = -1
if exp.Valid {
r.ExpectedChecks = int(exp.Int64)
}
out = append(out, r)
}
return out, rows.Err()
}
// RunCheck is one stored check of a run's result cycle, with what the
// analytics needs to place it.
type RunCheck struct {
RegistryID int64
Source string
CheckType string
Target string
Success bool
ValidatorID string
Detail string
RecordedAt time.Time
AfterVerdict bool
}
// EachRunCheck calls fn for every check of the cycles that make up the run's
// results (the latest cycle of each address in the run), in one pass over the
// run_id index. fn must not call back into the DB (one connection).
func (d *DB) EachRunCheck(ctx context.Context, runID int64, fn func(RunCheck)) error {
rows, err := d.QueryContext(ctx, `
SELECT c.registry_id, c.source, c.check_type, c.target, c.success, c.validator_id, c.detail,
COALESCE(c.recorded_at, c.created_at), c.after_verdict
FROM checks c JOIN run_results r ON r.run_id=c.run_id AND r.registry_id=c.registry_id AND r.cycle_id=c.cycle_id
WHERE c.run_id=?`, runID)
if err != nil {
return err
}
defer rows.Close()
for rows.Next() {
var c RunCheck
var rec string
if err := rows.Scan(&c.RegistryID, &c.Source, &c.CheckType, &c.Target, &c.Success, &c.ValidatorID, &c.Detail, &rec, &c.AfterVerdict); err != nil {
return err
}
if c.RecordedAt, err = dbToTime(rec); err != nil {
return err
}
fn(c)
}
return rows.Err()
}
// RunDataVersion changes whenever a check of the run is written, so a cache of
// computed analytics can tell when it is stale.
func (d *DB) RunDataVersion(ctx context.Context, runID int64) (string, error) {
var maxID, n sql.NullInt64
var rec sql.NullString
if err := d.QueryRowContext(ctx, `SELECT MAX(id), COUNT(*), MAX(recorded_at) FROM checks WHERE run_id=?`, runID).Scan(&maxID, &n, &rec); err != nil {
return "", err
}
var res sql.NullString
if err := d.QueryRowContext(ctx, `SELECT MAX(aggregated_at) FROM run_results WHERE run_id=?`, runID).Scan(&res); err != nil {
return "", err
}
return fmt.Sprintf("%d/%d/%s/%s", maxID.Int64, n.Int64, rec.String, res.String), nil
}
// Subnet is one entry of the administrator's subnet list.
type Subnet struct {
CIDR string
Label string
}
// ListSubnets returns the configured subnets, sorted by prefix.
func (d *DB) ListSubnets(ctx context.Context) ([]Subnet, error) {
rows, err := d.QueryContext(ctx, `SELECT cidr, label FROM subnets`)
if err != nil {
return nil, err
}
defer rows.Close()
var out []Subnet
for rows.Next() {
var s Subnet
if err := rows.Scan(&s.CIDR, &s.Label); err != nil {
return nil, err
}
out = append(out, s)
}
if err := rows.Err(); err != nil {
return nil, err
}
sort.Slice(out, func(i, j int) bool {
a, _ := netip.ParsePrefix(out[i].CIDR)
b, _ := netip.ParsePrefix(out[j].CIDR)
if a.Addr() != b.Addr() {
return a.Addr().Less(b.Addr())
}
return a.Bits() < b.Bits()
})
return out, nil
}
// ReplaceSubnets replaces the whole subnet list. Every entry must be a valid
// CIDR; entries are stored in canonical form (host bits cleared) and
// duplicates collapse.
func (d *DB) ReplaceSubnets(ctx context.Context, subnets []Subnet) error {
canon := map[string]string{}
for _, s := range subnets {
p, err := netip.ParsePrefix(s.CIDR)
if err != nil {
return fmt.Errorf("subnet %q: %v: %w", s.CIDR, err, ErrValidation)
}
canon[p.Masked().String()] = s.Label
}
tx, err := d.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if _, err := tx.ExecContext(ctx, `DELETE FROM subnets`); err != nil {
return err
}
for cidr, label := range canon {
if _, err := tx.ExecContext(ctx, `INSERT INTO subnets (cidr, label) VALUES (?, ?)`, cidr, label); err != nil {
return err
}
}
return tx.Commit()
}
// CountRecheckedInRun returns how many addresses have more than one cycle of
// checks inside the run (re-checked while the run was open).
func (d *DB) CountRecheckedInRun(ctx context.Context, runID int64) (int, error) {
var n int
err := d.QueryRowContext(ctx, `
SELECT COUNT(*) FROM (SELECT registry_id FROM checks WHERE run_id=? GROUP BY registry_id HAVING COUNT(DISTINCT cycle_id) > 1)
`, runID).Scan(&n)
return n, err
}