package db import ( "context" "database/sql" "path/filepath" "testing" "time" ) // checkingIP submits one address and moves it to the checking state. func checkingIP(t *testing.T, d *DB, addr string) *IPQueueItem { t.Helper() ctx := context.Background() if _, err := d.SubmitIPs(ctx, []string{addr}); err != nil { t.Fatal(err) } ip, err := d.GetIPByAddress(ctx, addr) if err != nil { t.Fatal(err) } if err := d.SetChecking(ctx, ip.ID, time.Minute); err != nil { t.Fatal(err) } ip, err = d.GetIP(ctx, ip.ID) if err != nil { t.Fatal(err) } return ip } func checkOf(ip *IPQueueItem, ct string, ok bool) Check { return Check{IPID: ip.ID, IPAddress: ip.IPAddress, AttemptNumber: ip.AttemptNumber, Source: InboundSource(1), CheckType: ct, Target: ip.IPAddress, Success: ok, CheckedAt: Now()} } func storedSuccess(t *testing.T, d *DB, ip *IPQueueItem, ct string) (success bool, found bool) { t.Helper() err := d.QueryRowContext(context.Background(), `SELECT success FROM checks WHERE ip_id=? AND check_type=?`, ip.ID, ct).Scan(&success) if err == sql.ErrNoRows { return false, false } if err != nil { t.Fatal(err) } return success, true } // Results are accepted while the address is being checked (also repeated ones, // idempotently) and refused from the moment the verdict is being computed. func TestUpsertCheckIfOpenFreezesAtVerdict(t *testing.T) { d, ctx := newTestDB(t) ip := checkingIP(t, d, "1.2.3.4") for i := 0; i < 2; i++ { if ok, err := d.UpsertCheckIfOpen(ctx, checkOf(ip, "tcp-22", true)); err != nil || !ok { t.Fatalf("write %d while checking: ok=%v err=%v", i, ok, err) } } var n int if err := d.QueryRowContext(ctx, `SELECT COUNT(*) FROM checks WHERE ip_id=?`, ip.ID).Scan(&n); err != nil || n != 1 { t.Fatalf("expected one row after repeated write, got %d err=%v", n, err) } // An older attempt's result does not touch the current attempt. old := checkOf(ip, "icmp", true) old.AttemptNumber = ip.AttemptNumber - 1 if ok, err := d.UpsertCheckIfOpen(ctx, old); err != nil || ok { t.Fatalf("stale attempt must be dropped: ok=%v err=%v", ok, err) } if _, found := storedSuccess(t, d, ip, "icmp"); found { t.Fatal("stale attempt wrote a row") } if err := d.SetAggregating(ctx, ip.ID); err != nil { t.Fatal(err) } // Neither a new check nor a change of an existing one gets through. if ok, err := d.UpsertCheckIfOpen(ctx, checkOf(ip, "ssh", false)); err != nil || ok { t.Fatalf("new check while aggregating must be dropped: ok=%v err=%v", ok, err) } if ok, err := d.UpsertCheckIfOpen(ctx, checkOf(ip, "tcp-22", false)); err != nil || ok { t.Fatalf("overwrite while aggregating must be dropped: ok=%v err=%v", ok, err) } if err := d.FinishIP(ctx, ip.ID, ResultPass); err != nil { t.Fatal(err) } if ok, err := d.UpsertCheckIfOpen(ctx, checkOf(ip, "tcp-22", false)); err != nil || ok { t.Fatalf("overwrite after the verdict must be dropped: ok=%v err=%v", ok, err) } if s, found := storedSuccess(t, d, ip, "tcp-22"); !found || !s { t.Fatalf("stored tcp-22 changed after the verdict: found=%v success=%v", found, s) } if _, found := storedSuccess(t, d, ip, "ssh"); found { t.Fatal("a check written after the verdict") } } // recorded_at is the server's own write time and moves on every accepted write. func TestUpsertCheckIfOpenSetsRecordedAt(t *testing.T) { d, ctx := newTestDB(t) ip := checkingIP(t, d, "1.2.3.4") if _, err := d.UpsertCheckIfOpen(ctx, checkOf(ip, "icmp", true)); err != nil { t.Fatal(err) } var first string if err := d.QueryRowContext(ctx, `SELECT recorded_at FROM checks WHERE ip_id=?`, ip.ID).Scan(&first); err != nil || first == "" { t.Fatalf("recorded_at not set: %q %v", first, err) } time.Sleep(5 * time.Millisecond) if _, err := d.UpsertCheckIfOpen(ctx, checkOf(ip, "icmp", false)); err != nil { t.Fatal(err) } var second, created string if err := d.QueryRowContext(ctx, `SELECT recorded_at, created_at FROM checks WHERE ip_id=?`, ip.ID).Scan(&second, &created); err != nil { t.Fatal(err) } if second <= first || created != first { t.Fatalf("recorded_at must advance, created_at stay: first=%s second=%s created=%s", first, second, created) } } // A prober site is handed an address until it reports it complete in the // current attempt; other sites are unaffected; a new attempt hands it out again. func TestListCheckingForSite(t *testing.T) { d, ctx := newTestDB(t) a := checkingIP(t, d, "1.1.1.1") b := checkingIP(t, d, "2.2.2.2") names := func(items []IPQueueItem) string { s := "" for _, it := range items { s += it.IPAddress + " " } return s } for site := 1; site <= 2; site++ { items, err := d.ListCheckingForSite(ctx, site) if err != nil || len(items) != 2 { t.Fatalf("site %d before any report: %q err=%v", site, names(items), err) } } if err := d.SetSiteComplete(ctx, a.ID, 1); err != nil { t.Fatal(err) } if items, _ := d.ListCheckingForSite(ctx, 1); names(items) != "2.2.2.2 " { t.Fatalf("site 1 after completing 1.1.1.1: %q", names(items)) } if items, _ := d.ListCheckingForSite(ctx, 2); len(items) != 2 { t.Fatalf("site 2 must still get both: %q", names(items)) } // A retry starts a new attempt: site 1 has to probe the address again. if err := d.RequeueOrFail(ctx, a.ID, "", 3); err != nil { t.Fatal(err) } if err := d.SetChecking(ctx, a.ID, time.Minute); err != nil { t.Fatal(err) } if items, _ := d.ListCheckingForSite(ctx, 1); len(items) != 2 { t.Fatalf("site 1 after a new attempt: %q", names(items)) } _ = b } // The aggregation window counts from the start of checking; a retry clears it. func TestCheckingStartedAt(t *testing.T) { d, ctx := newTestDB(t) ip := checkingIP(t, d, "1.2.3.4") if ip.CheckingStartedAt == nil || time.Since(*ip.CheckingStartedAt) > time.Minute { t.Fatalf("checking_started_at not set: %v", ip.CheckingStartedAt) } if err := d.RequeueOrFail(ctx, ip.ID, "", 3); err != nil { t.Fatal(err) } ip, err := d.GetIP(ctx, ip.ID) if err != nil { t.Fatal(err) } if ip.CheckingStartedAt != nil { t.Fatalf("checking_started_at must be cleared on requeue, got %v", ip.CheckingStartedAt) } } // Migration 0010 flags the rows of an existing database that were written // after their address's verdict, and leaves the others alone. func TestMigration0010MarksRowsAfterVerdict(t *testing.T) { ctx := context.Background() path := filepath.Join(t.TempDir(), "old.db") raw, err := sql.Open("sqlite", path) if err != nil { t.Fatal(err) } raw.SetMaxOpenConns(1) for _, m := range migrations { if m.version > 9 { break } if _, err := raw.ExecContext(ctx, m.sql); err != nil { t.Fatalf("migration %d: %v", m.version, err) } } if _, err := raw.ExecContext(ctx, `PRAGMA user_version=9`); err != nil { t.Fatal(err) } for _, q := range []string{ `INSERT INTO ip_registry (id, ip_address, first_seen_at, last_seen_at, next_cycle, created_at, updated_at) VALUES (1, '1.2.3.4', '2026-10-02T13:00:00Z', '2026-10-02T13:00:00Z', 2, '2026-10-02T13:00:00Z', '2026-10-02T13:00:00Z')`, `INSERT INTO ip_queue (id, ip_address, sequence, state, overall_result, aggregated_at, registry_id, cycle_id, created_at, updated_at) VALUES (1, '1.2.3.4', 1, 'done', 'pass', '2026-10-02T13:48:45.659Z', 1, 1, '2026-10-02T13:00:00Z', '2026-10-02T13:00:00Z')`, // before the verdict, and (with a longer fraction) after it `INSERT INTO checks (registry_id, cycle_id, ip_id, ip_address, attempt_number, validator_id, source, check_type, target, success, checked_at, created_at) VALUES (1, 1, 1, '1.2.3.4', 1, '', 'inbound-site-1', 'icmp', '1.2.3.4', 1, '2026-10-02T13:48:40.100000000Z', '2026-10-02T13:48:40.2Z')`, `INSERT INTO checks (registry_id, cycle_id, ip_id, ip_address, attempt_number, validator_id, source, check_type, target, success, checked_at, created_at) VALUES (1, 1, 1, '1.2.3.4', 1, '', 'inbound-site-1', 'ssh', '1.2.3.4', 0, '2026-10-02T13:48:45.730314288Z', '2026-10-02T13:48:40.3Z')`, } { if _, err := raw.ExecContext(ctx, q); err != nil { t.Fatalf("seed: %v\n%s", err, q) } } raw.Close() d, err := Open(ctx, path) if err != nil { t.Fatalf("open (runs migration 10): %v", err) } defer d.Close() flag := func(ct string) (after int, recorded, created string) { if err := d.QueryRowContext(ctx, `SELECT after_verdict, recorded_at, created_at FROM checks WHERE check_type=?`, ct).Scan(&after, &recorded, &created); err != nil { t.Fatal(err) } return } if a, rec, cr := flag("icmp"); a != 0 || rec != cr { t.Errorf("icmp: after_verdict=%d recorded=%s created=%s", a, rec, cr) } if a, rec, cr := flag("ssh"); a != 1 || rec != cr { t.Errorf("ssh: after_verdict=%d recorded=%s created=%s", a, rec, cr) } var ver int if err := d.QueryRowContext(ctx, `PRAGMA user_version`).Scan(&ver); err != nil || ver != 11 { t.Errorf("user_version=%d err=%v", ver, err) } }