387 lines
12 KiB
Go
387 lines
12 KiB
Go
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
|
||
|
|
}
|