package db import ( "context" "database/sql" "fmt" "net/netip" "sort" "strings" "time" ) // findOrCreateRegistryTx returns the ip_registry row id for addr, creating // it (next_cycle starting at 1) the first time this address is ever seen. // Safe to call on every (re)submission of an address — ip_address is // UNIQUE, so a second call for the same address is a no-op lookup. func findOrCreateRegistryTx(ctx context.Context, tx *sql.Tx, addr string, now string) (int64, error) { if _, err := tx.ExecContext(ctx, ` INSERT INTO ip_registry (ip_address, first_seen_at, last_seen_at, next_cycle, created_at, updated_at) VALUES (?, ?, ?, 1, ?, ?) ON CONFLICT(ip_address) DO NOTHING `, addr, now, now, now, now); err != nil { return 0, fmt.Errorf("find or create registry for %s: %w", addr, err) } var id int64 if err := tx.QueryRowContext(ctx, `SELECT id FROM ip_registry WHERE ip_address=?`, addr).Scan(&id); err != nil { return 0, fmt.Errorf("lookup registry id for %s: %w", addr, err) } return id, nil } // nextRegistryCycleTx allocates the next cycle_id for registryID and // advances last_seen_at. Call once per fresh generation of check history for // an address — a brand new ip_queue row, an admin-triggered requeue, or a // system retry that bumps attempt_number — so every generation's checks get // a cycle_id no other generation, past or future (even across ip_queue row // deletion and recreation), will ever reuse. func nextRegistryCycleTx(ctx context.Context, tx *sql.Tx, registryID int64, now string) (int, error) { var cycle int if err := tx.QueryRowContext(ctx, `SELECT next_cycle FROM ip_registry WHERE id=?`, registryID).Scan(&cycle); err != nil { return 0, fmt.Errorf("read next_cycle for registry %d: %w", registryID, err) } if _, err := tx.ExecContext(ctx, ` UPDATE ip_registry SET next_cycle=next_cycle+1, last_seen_at=?, updated_at=? WHERE id=? `, now, now, registryID); err != nil { return 0, fmt.Errorf("advance next_cycle for registry %d: %w", registryID, err) } return cycle, nil } // RegistrySummary is one row of the full IP registry, combining the durable // ip_registry record with a rollup of its check history and, if the address // currently has a live ip_queue row, that row's state. type RegistrySummary struct { RegistryItem TotalCycles int LastResult string LastCheckedAt *time.Time InQueue bool CurrentState string // LastCycleID is the newest cycle_id with recorded checks (0 if none). // Egress and Ingress count that cycle's recorded checks per level. LastCycleID int Egress LevelResult Ingress LevelResult } // TypeStat counts the recorded checks of one check family (see CheckFamily) // within a level: Total checks, OK of them successful. type TypeStat struct { Type string Total int OK int } // LevelResult is the "OK of Total" rollup of one level (egress or ingress) of // a cycle, with the same counts split by check family, sorted by Type. type LevelResult struct { Total int OK int ByType []TypeStat } // ListRegistry returns every address ever submitted, newest first-seen // last, each with a summary of its accumulated check history. Reads two // simple queries plus one aggregate rather than a single large join, since // the registry is expected to stay small enough (one row per distinct // address ever seen) that this is simpler to reason about than a // multi-way correlated subquery. func (d *DB) ListRegistry(ctx context.Context) ([]RegistrySummary, error) { rows, err := d.QueryContext(ctx, ` SELECT id, ip_address, first_seen_at, last_seen_at, next_cycle, created_at, updated_at FROM ip_registry ORDER BY first_seen_at, id `) if err != nil { return nil, err } defer rows.Close() items, err := scanRegistryItems(rows) if err != nil { return nil, err } out := make([]RegistrySummary, len(items)) for i, item := range items { s := RegistrySummary{RegistryItem: item} if err := d.fillRegistrySummary(ctx, &s); err != nil { return nil, err } out[i] = s } if err := d.fillRegistryLevels(ctx, out); err != nil { return nil, err } return out, nil } // RegistryFilter narrows ListRegistryPage. The zero value matches everything. type RegistryFilter struct { Query string // substring of ip_address LastResult string // pass|partial|fail|cancelled — same meaning as RegistrySummary.LastResult RunID int64 // only addresses that have a result in this run Subnet string // only addresses inside this CIDR } // subnetIDs returns the registry ids of the addresses inside prefix. SQLite // has no CIDR operators, so the registry's addresses are filtered here. func (d *DB) subnetIDs(ctx context.Context, cidr string) ([]any, error) { p, err := netip.ParsePrefix(cidr) if err != nil { return nil, fmt.Errorf("subnet %q: %v: %w", cidr, err, ErrValidation) } rows, err := d.QueryContext(ctx, `SELECT id, ip_address FROM ip_registry`) if err != nil { return nil, err } defer rows.Close() var ids []any for rows.Next() { var id int64 var ip string if err := rows.Scan(&id, &ip); err != nil { return nil, err } if a, err := netip.ParseAddr(ip); err == nil && p.Contains(a) { ids = append(ids, id) } } return ids, rows.Err() } // lastResultCond is the SQL form of fillRegistrySummary's LastResult rule, // over `ip_registry r LEFT JOIN ip_queue q ON q.registry_id=r.id`: a live // queue row's overall_result is authoritative (and a live row without one has // no verdict); with no live row the latest recorded cycle is classified from // its checks (pass if all succeeded, fail if none, partial otherwise). const lastResultCond = `( (q.id IS NOT NULL AND q.overall_result = ?) OR (q.id IS NULL AND ( SELECT CASE WHEN SUM(c.success) = 0 THEN 'fail' WHEN SUM(c.success) = COUNT(*) THEN 'pass' ELSE 'partial' END FROM checks c WHERE c.registry_id = r.id AND c.cycle_id = (SELECT MAX(c2.cycle_id) FROM checks c2 WHERE c2.registry_id = r.id) HAVING COUNT(*) > 0 ) = ?) )` // ListRegistryPage returns one page (limit/offset) of the registry in the same // order as ListRegistry, filtered by f, plus the total number of matching // rows. LIMIT/OFFSET are applied in SQL before the per-row summary queries, // so only the rows of the page pay for them. limit <= 0 means no limit. func (d *DB) ListRegistryPage(ctx context.Context, f RegistryFilter, limit, offset int) ([]RegistrySummary, int, error) { var conds []string var args []any if f.Query != "" { conds = append(conds, "instr(r.ip_address, ?) > 0") args = append(args, f.Query) } if f.LastResult != "" { conds = append(conds, lastResultCond) args = append(args, f.LastResult, f.LastResult) } if f.RunID > 0 { conds = append(conds, "r.id IN (SELECT registry_id FROM run_results WHERE run_id = ?)") args = append(args, f.RunID) } if f.Subnet != "" { ids, err := d.subnetIDs(ctx, f.Subnet) if err != nil { return nil, 0, err } if len(ids) == 0 { conds = append(conds, "0 = 1") } else { conds = append(conds, "r.id IN ("+strings.TrimSuffix(strings.Repeat("?,", len(ids)), ",")+")") args = append(args, ids...) } } from := ` FROM ip_registry r LEFT JOIN ip_queue q ON q.registry_id = r.id ` where := "" if len(conds) > 0 { where = "WHERE " + strings.Join(conds, " AND ") + " " } var total int if err := d.QueryRowContext(ctx, `SELECT COUNT(*)`+from+where, args...).Scan(&total); err != nil { return nil, 0, err } q := `SELECT r.id, r.ip_address, r.first_seen_at, r.last_seen_at, r.next_cycle, r.created_at, r.updated_at` + from + where + `ORDER BY r.first_seen_at, r.id` qargs := append([]any(nil), args...) if limit > 0 { q += " LIMIT ? OFFSET ?" qargs = append(qargs, limit, offset) } rows, err := d.QueryContext(ctx, q, qargs...) if err != nil { return nil, 0, err } items, err := scanRegistryItems(rows) rows.Close() if err != nil { return nil, 0, err } out := make([]RegistrySummary, len(items)) for i, item := range items { s := RegistrySummary{RegistryItem: item} if err := d.fillRegistrySummary(ctx, &s); err != nil { return nil, 0, err } out[i] = s } if err := d.fillRegistryLevels(ctx, out); err != nil { return nil, 0, err } return out, total, nil } // GetRegistryByAddress returns the registry row (with summary) for a single // address, or ErrNotFound if it has never been submitted. func (d *DB) GetRegistryByAddress(ctx context.Context, address string) (*RegistrySummary, error) { row := d.QueryRowContext(ctx, ` SELECT id, ip_address, first_seen_at, last_seen_at, next_cycle, created_at, updated_at FROM ip_registry WHERE ip_address=? `, address) item, err := scanRegistryItem(row) if err != nil { if err == sql.ErrNoRows { return nil, fmt.Errorf("ip %q: %w", address, ErrNotFound) } return nil, err } s := &RegistrySummary{RegistryItem: *item} if err := d.fillRegistrySummary(ctx, s); err != nil { return nil, err } one := []RegistrySummary{*s} if err := d.fillRegistryLevels(ctx, one); err != nil { return nil, err } *s = one[0] return s, nil } // fillRegistrySummary computes the badge-level summary shown on the // registry list/detail pages. LastResult must reflect the *aggregated* // outcome of the most recent cycle (pass/partial/fail/cancelled), not the // success flag of whichever individual check happens to have the latest // checked_at — a cycle with a mix of passing and failing checks (e.g. one // egress target timed out while the rest succeeded) is "partial", even // though the chronologically-last check to report in might have passed. func (d *DB) fillRegistrySummary(ctx context.Context, s *RegistrySummary) error { if err := d.QueryRowContext(ctx, ` SELECT COUNT(DISTINCT cycle_id) FROM checks WHERE registry_id=? `, s.ID).Scan(&s.TotalCycles); err != nil { return err } var lastCheckedAt sql.NullString if err := d.QueryRowContext(ctx, ` SELECT checked_at FROM checks WHERE registry_id=? ORDER BY cycle_id DESC, checked_at DESC LIMIT 1 `, s.ID).Scan(&lastCheckedAt); err != nil && err != sql.ErrNoRows { return err } t, err := nullStringToTimePtr(lastCheckedAt) if err != nil { return err } s.LastCheckedAt = t var state, overallResult sql.NullString err = d.QueryRowContext(ctx, `SELECT state, overall_result FROM ip_queue WHERE registry_id=?`, s.ID). Scan(&state, &overallResult) if err != nil && err != sql.ErrNoRows { return err } if err == nil { s.InQueue = true s.CurrentState = state.String } switch { case overallResult.Valid && overallResult.String != "": // The address has a live ip_queue row with a finished cycle // (done/failed) — overall_result is the orchestrator's own // aggregation (internal/orchestrator.aggregateAndRelease / // db.CancelIP), the authoritative source of truth. Use it as-is // rather than re-deriving it from raw check rows. s.LastResult = overallResult.String case s.InQueue: // A live ip_queue row exists but its current cycle hasn't finished // yet (still queued/checking/etc, overall_result not set) — no // verdict to show yet; leave LastResult empty rather than guessing // from a still-incomplete set of checks. default: // No live ip_queue row (deleted from the queue) — fall back to // classifying the most recent recorded cycle from its own checks, // the same pass/fail/partial rule aggregateAndRelease uses (minus // "missing" checks, which aren't knowable after the fact). result, err := lastCycleResultFromChecks(ctx, d, s.ID) if err != nil { return err } s.LastResult = result } return nil } // registryLevelsChunk bounds the number of registry ids per query, well under // SQLite's bound-variable limit. const registryLevelsChunk = 500 // registryLevelsQuery is the grouped query of fillRegistryLevels for n // registry ids: per address, the counts of its newest cycle by source and // check type. func registryLevelsQuery(n int) string { return ` SELECT c.registry_id, m.cid, c.source, c.check_type, COUNT(*), COALESCE(SUM(c.success), 0) FROM checks c JOIN (SELECT registry_id, MAX(cycle_id) AS cid FROM checks WHERE registry_id IN (` + strings.TrimSuffix(strings.Repeat("?,", n), ",") + `) GROUP BY registry_id) m ON m.registry_id = c.registry_id AND m.cid = c.cycle_id GROUP BY c.registry_id, m.cid, c.source, c.check_type ` } // fillRegistryLevels sets LastCycleID, Egress and Ingress on every summary in // sums: the recorded checks of each address's newest cycle (the same cycle // whose time is LastCheckedAt), counted per level and per check family. It // runs one grouped query per chunk of addresses, not one per address, over // idx_checks_registry_cycle. Checks whose source is neither egress nor an // inbound site are not counted. The counts follow the recorded rows only, so // they can differ from LastResult, which also treats missing results as // failures. func (d *DB) fillRegistryLevels(ctx context.Context, sums []RegistrySummary) error { pos := make(map[int64]int, len(sums)) for i := range sums { pos[sums[i].ID] = i } type key struct { id int64 level, family string } for start := 0; start < len(sums); start += registryLevelsChunk { end := min(start+registryLevelsChunk, len(sums)) args := make([]any, 0, end-start) for _, s := range sums[start:end] { args = append(args, s.ID) } rows, err := d.QueryContext(ctx, registryLevelsQuery(len(args)), args...) if err != nil { return err } stats := map[key]*TypeStat{} for rows.Next() { var id int64 var cid, total, ok int var source, checkType string if err := rows.Scan(&id, &cid, &source, &checkType, &total, &ok); err != nil { rows.Close() return err } sums[pos[id]].LastCycleID = cid level := CheckLevel(source) if level == "" { continue } k := key{id, level, CheckFamily(checkType)} st := stats[k] if st == nil { st = &TypeStat{Type: k.family} stats[k] = st } st.Total += total st.OK += ok } err = rows.Err() rows.Close() if err != nil { return err } for k, st := range stats { lr := &sums[pos[k.id]].Egress if k.level == LevelIngress { lr = &sums[pos[k.id]].Ingress } lr.Total += st.Total lr.OK += st.OK lr.ByType = append(lr.ByType, *st) } } for i := range sums { for _, lr := range []*LevelResult{&sums[i].Egress, &sums[i].Ingress} { sort.Slice(lr.ByType, func(a, b int) bool { return lr.ByType[a].Type < lr.ByType[b].Type }) } } return nil } // lastCycleResultFromChecks classifies the most recent cycle recorded for // registryID directly from its checks rows: pass if every recorded check // succeeded, fail if every one failed, partial on a mix. Returns "" if no // checks are recorded at all. func lastCycleResultFromChecks(ctx context.Context, d *DB, registryID int64) (string, error) { var cycle sql.NullInt64 if err := d.QueryRowContext(ctx, `SELECT MAX(cycle_id) FROM checks WHERE registry_id=?`, registryID).Scan(&cycle); err != nil { return "", err } if !cycle.Valid { return "", nil } var total, passed int if err := d.QueryRowContext(ctx, ` SELECT COUNT(*), COALESCE(SUM(success), 0) FROM checks WHERE registry_id=? AND cycle_id=? `, registryID, cycle.Int64).Scan(&total, &passed); err != nil { return "", err } switch { case total == 0: return "", nil case passed == 0: return ResultFail, nil case passed == total: return ResultPass, nil default: return ResultPartial, nil } } func scanRegistryItems(rows *sql.Rows) ([]RegistryItem, error) { var out []RegistryItem for rows.Next() { item, err := scanRegistryItem(rows) if err != nil { return nil, err } out = append(out, *item) } return out, rows.Err() } func scanRegistryItem(row rowScanner) (*RegistryItem, error) { var item RegistryItem var firstSeenAt, lastSeenAt, createdAt, updatedAt string if err := row.Scan(&item.ID, &item.IPAddress, &firstSeenAt, &lastSeenAt, &item.NextCycle, &createdAt, &updatedAt); err != nil { return nil, err } var err error if item.FirstSeenAt, err = dbToTime(firstSeenAt); err != nil { return nil, err } if item.LastSeenAt, err = dbToTime(lastSeenAt); err != nil { return nil, err } if item.CreatedAt, err = dbToTime(createdAt); err != nil { return nil, err } if item.UpdatedAt, err = dbToTime(updatedAt); err != nil { return nil, err } return &item, nil } // PruneRegistryHistory deletes all but the newest keepCycles cycles' worth // of check history for registryID — a no-op if keepCycles <= 0 (unlimited // retention) or the address has that many cycles or fewer. events rows tied // to this registry are pruned to the same cutoff; ip_registry itself and its // next_cycle counter are never touched, so pruning never risks a future // cycle_id collision. func (d *DB) PruneRegistryHistory(ctx context.Context, registryID int64, keepCycles int) error { if keepCycles <= 0 { return nil } var cutoff sql.NullInt64 err := d.QueryRowContext(ctx, ` SELECT MIN(cycle_id) FROM ( SELECT DISTINCT cycle_id FROM checks WHERE registry_id=? ORDER BY cycle_id DESC LIMIT ? ) `, registryID, keepCycles).Scan(&cutoff) if err != nil { return fmt.Errorf("find prune cutoff for registry %d: %w", registryID, err) } if !cutoff.Valid { return nil // fewer than keepCycles cycles recorded — nothing to prune } if _, err := d.ExecContext(ctx, `DELETE FROM checks WHERE registry_id=? AND cycle_id