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 }