Files
routewatch/internal/database/database_test.go
T
sneak cd59cb8a8d
check / check (push) Failing after 0s
Serve /api/v1/stats from a cache and index-scan the route timestamps (closes #27)
Once the database passed about 4.5 GiB every stats request ran a COUNT(*)
over each table plus a MIN/MAX union scan of both route tables, took the
full timeout and returned HTTP 500, so the status page went blank.

The server now keeps the last database statistics in memory and recomputes
them at most once every 30 seconds; requests serve the cached copy and a
stale copy triggers a single background refresh, so no request runs the
scans. The route-count split is folded into the cached stats, removing the
separate per-request live-route count query.

The oldest/newest route timestamps now read one row from each end of the
last_updated index instead of scanning both tables, and select the column
directly so the driver parses it into time.Time; the old aggregate returned
an untyped string that failed to scan and logged a warning every call.

Model: opus-4-8
2026-09-21 23:32:11 +00:00

545 lines
14 KiB
Go

package database
import (
"context"
"database/sql"
"net"
"sync"
"testing"
"time"
"git.eeqj.de/sneak/routewatch/internal/config"
"git.eeqj.de/sneak/routewatch/internal/logger"
"github.com/google/uuid"
)
// tempStoreMemory is the PRAGMA temp_store value meaning "hold temp B-trees in
// memory"; the DSN change must leave temp_store below this so they spill to disk.
const tempStoreMemory = 2
// heldConnections is how many pooled connections the pragma test holds open at
// once so each is a distinct SQLite connection that parsed the DSN.
const heldConnections = 5
// Parameters for the checkpoint-contention regression test.
const (
// contentionIterations is how many batch writes race the checkpoint loop.
contentionIterations = 400
// contendedASNCount is the small set of ASNs the batches reuse, so most
// batches update existing rows and exercise the read-before-write path.
contendedASNCount = 16
// asnSecondBand offsets a second ASN per batch so each batch writes more
// than one row.
asnSecondBand = 100
)
func TestIPToUint32(t *testing.T) {
tests := []struct {
name string
ip string
expected uint32
}{
{
name: "Simple IP",
ip: "192.168.1.1",
expected: 3232235777, // 192<<24 + 168<<16 + 1<<8 + 1
},
{
name: "Minimum IP",
ip: "0.0.0.0",
expected: 0,
},
{
name: "Maximum IP",
ip: "255.255.255.255",
expected: 4294967295,
},
{
name: "10.0.0.0",
ip: "10.0.0.0",
expected: 167772160,
},
{
name: "172.16.0.0",
ip: "172.16.0.0",
expected: 2886729728,
},
{
name: "8.8.8.8",
ip: "8.8.8.8",
expected: 134744072,
},
{
name: "1.2.3.4",
ip: "1.2.3.4",
expected: 16909060,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
ip := net.ParseIP(tt.ip)
if ip == nil {
t.Fatalf("Failed to parse IP: %s", tt.ip)
}
result := ipToUint32(ip)
if result != tt.expected {
t.Errorf("ipToUint32(%s) = %d, want %d", tt.ip, result, tt.expected)
}
// Test with IPv4-mapped IPv6 address
ip6 := net.ParseIP(tt.ip).To16()
if ip6 != nil {
result6 := ipToUint32(ip6)
if result6 != tt.expected {
t.Errorf("ipToUint32(%s as IPv6) = %d, want %d", tt.ip, result6, tt.expected)
}
}
})
}
}
func TestCalculateIPv4Range(t *testing.T) {
tests := []struct {
name string
cidr string
wantStart uint32
wantEnd uint32
wantErr bool
}{
{
name: "Single IP /32",
cidr: "192.168.1.1/32",
wantStart: 3232235777,
wantEnd: 3232235777,
},
{
name: "Class C /24",
cidr: "192.168.1.0/24",
wantStart: 3232235776, // 192.168.1.0
wantEnd: 3232236031, // 192.168.1.255
},
{
name: "Class B /16",
cidr: "192.168.0.0/16",
wantStart: 3232235520, // 192.168.0.0
wantEnd: 3232301055, // 192.168.255.255
},
{
name: "Class A /8",
cidr: "10.0.0.0/8",
wantStart: 167772160, // 10.0.0.0
wantEnd: 184549375, // 10.255.255.255
},
{
name: "Entire IPv4 space /0",
cidr: "0.0.0.0/0",
wantStart: 0,
wantEnd: 4294967295,
},
{
name: "Small subnet /30",
cidr: "192.168.1.0/30",
wantStart: 3232235776, // 192.168.1.0
wantEnd: 3232235779, // 192.168.1.3
},
{
name: "Medium subnet /20",
cidr: "172.16.0.0/20",
wantStart: 2886729728, // 172.16.0.0
wantEnd: 2886733823, // 172.16.15.255
},
{
name: "Private range 172.16/12",
cidr: "172.16.0.0/12",
wantStart: 2886729728, // 172.16.0.0
wantEnd: 2887778303, // 172.31.255.255
},
{
name: "Google DNS /29",
cidr: "8.8.8.8/29",
wantStart: 134744072, // 8.8.8.8 (network is actually 8.8.8.8 with /29)
wantEnd: 134744079, // 8.8.8.15
},
{
name: "Non-zero host bits",
cidr: "192.168.1.5/24",
wantStart: 3232235776, // 192.168.1.0 (network address)
wantEnd: 3232236031, // 192.168.1.255
},
{
name: "Invalid CIDR",
cidr: "192.168.1.1/33",
wantErr: true,
},
{
name: "Invalid IP",
cidr: "256.256.256.256/24",
wantErr: true,
},
{
name: "IPv6 CIDR",
cidr: "2001:db8::/32",
wantErr: true,
},
{
name: "Empty CIDR",
cidr: "",
wantErr: true,
},
{
name: "Missing mask",
cidr: "192.168.1.1",
wantErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
start, end, err := CalculateIPv4Range(tt.cidr)
if tt.wantErr {
if err == nil {
t.Errorf("CalculateIPv4Range(%s) expected error, got nil", tt.cidr)
}
return
}
if err != nil {
t.Errorf("CalculateIPv4Range(%s) unexpected error: %v", tt.cidr, err)
return
}
if start != tt.wantStart {
t.Errorf("CalculateIPv4Range(%s) start = %d, want %d", tt.cidr, start, tt.wantStart)
}
if end != tt.wantEnd {
t.Errorf("CalculateIPv4Range(%s) end = %d, want %d", tt.cidr, end, tt.wantEnd)
}
// Verify that start <= end
if start > end {
t.Errorf("CalculateIPv4Range(%s) start (%d) > end (%d)", tt.cidr, start, end)
}
// Verify the range size matches the CIDR mask
if !tt.wantErr && tt.cidr != "" {
_, ipNet, _ := net.ParseCIDR(tt.cidr)
if ipNet != nil {
ones, bits := ipNet.Mask.Size()
expectedSize := uint32(1) << uint(bits-ones)
actualSize := end - start + 1
if actualSize != expectedSize {
t.Errorf("CalculateIPv4Range(%s) range size = %d, want %d", tt.cidr, actualSize, expectedSize)
}
}
}
})
}
}
func TestIPv4RangeIntegration(t *testing.T) {
// Test that our functions work correctly together
tests := []struct {
name string
cidr string
testIPs []string
shouldContain []bool
}{
{
name: "192.168.1.0/24",
cidr: "192.168.1.0/24",
testIPs: []string{
"192.168.1.0",
"192.168.1.1",
"192.168.1.255",
"192.168.0.255",
"192.168.2.0",
},
shouldContain: []bool{true, true, true, false, false},
},
{
name: "10.0.0.0/8",
cidr: "10.0.0.0/8",
testIPs: []string{
"10.0.0.0",
"10.255.255.255",
"10.1.2.3",
"9.255.255.255",
"11.0.0.0",
},
shouldContain: []bool{true, true, true, false, false},
},
{
name: "172.16.0.0/12",
cidr: "172.16.0.0/12",
testIPs: []string{
"172.16.0.0",
"172.31.255.255",
"172.20.1.1",
"172.15.255.255",
"172.32.0.0",
},
shouldContain: []bool{true, true, true, false, false},
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
start, end, err := CalculateIPv4Range(tt.cidr)
if err != nil {
t.Fatalf("Failed to calculate range for %s: %v", tt.cidr, err)
}
for i, testIP := range tt.testIPs {
ip := net.ParseIP(testIP)
if ip == nil {
t.Fatalf("Failed to parse test IP: %s", testIP)
}
ipUint := ipToUint32(ip)
contained := ipUint >= start && ipUint <= end
if contained != tt.shouldContain[i] {
t.Errorf("IP %s in range %s: got %v, want %v", testIP, tt.cidr, contained, tt.shouldContain[i])
}
}
})
}
}
// TestConnectionPoolPragmas holds several pooled connections open at once and
// checks each one carries the per-connection settings from the DSN, plus the
// process-wide hard heap limit.
func TestConnectionPoolPragmas(t *testing.T) {
cfg := &config.Config{StateDir: t.TempDir()}
db, err := New(cfg, logger.New())
if err != nil {
t.Fatalf("failed to create database: %v", err)
}
defer func() { _ = db.Close() }()
ctx := context.Background()
// Hold distinct connections open simultaneously so the pool must open a new
// one (each parsing the DSN) rather than hand back the same connection.
conns := make([]*sql.Conn, 0, heldConnections)
defer func() {
for _, c := range conns {
_ = c.Close()
}
}()
for i := 0; i < heldConnections; i++ {
c, err := db.db.Conn(ctx)
if err != nil {
t.Fatalf("failed to open connection %d: %v", i, err)
}
conns = append(conns, c)
}
for i, c := range conns {
var cacheSize int
if err := c.QueryRowContext(ctx, "PRAGMA cache_size").Scan(&cacheSize); err != nil {
t.Fatalf("conn %d: failed to read cache_size: %v", i, err)
}
if cacheSize != sqliteCacheSizeKiB {
t.Errorf("conn %d: cache_size = %d, want %d", i, cacheSize, sqliteCacheSizeKiB)
}
var busyTimeout int
if err := c.QueryRowContext(ctx, "PRAGMA busy_timeout").Scan(&busyTimeout); err != nil {
t.Fatalf("conn %d: failed to read busy_timeout: %v", i, err)
}
if busyTimeout != sqliteBusyTimeoutMs {
t.Errorf("conn %d: busy_timeout = %d, want %d", i, busyTimeout, sqliteBusyTimeoutMs)
}
var tempStore int
if err := c.QueryRowContext(ctx, "PRAGMA temp_store").Scan(&tempStore); err != nil {
t.Fatalf("conn %d: failed to read temp_store: %v", i, err)
}
if tempStore == tempStoreMemory {
t.Errorf("conn %d: temp_store = %d, want anything but %d (MEMORY)", i, tempStore, tempStoreMemory)
}
var hardHeapLimit int64
if err := c.QueryRowContext(ctx, "PRAGMA hard_heap_limit").Scan(&hardHeapLimit); err != nil {
t.Fatalf("conn %d: failed to read hard_heap_limit: %v", i, err)
}
if hardHeapLimit != sqliteHardHeapLimitBytes {
t.Errorf("conn %d: hard_heap_limit = %d, want %d", i, hardHeapLimit, sqliteHardHeapLimitBytes)
}
}
}
// TestBatchWriteDuringCheckpoint reproduces issue #25. A batch write reads
// (SELECT) before it writes (INSERT/UPDATE). Under the default deferred locking
// the transaction begins as a reader and, when the maintainer's WAL checkpoint
// holds the write lock, its upgrade to writer fails immediately with "database
// is locked" without honouring busy_timeout, dropping the batch. With
// _txlock=immediate the transaction takes the write lock at BEGIN and waits, so
// no batch is dropped. The checkpoint runs without the Database mutex, exactly
// as the background maintainer does in production.
func TestBatchWriteDuringCheckpoint(t *testing.T) {
cfg := &config.Config{StateDir: t.TempDir()}
db, err := New(cfg, logger.New())
if err != nil {
t.Fatalf("failed to create database: %v", err)
}
defer func() { _ = db.Close() }()
ctx, cancel := context.WithCancel(context.Background())
var wg sync.WaitGroup
wg.Add(1)
go func() {
defer wg.Done()
for {
select {
case <-ctx.Done():
return
default:
_ = db.Checkpoint(ctx) // errors are the checkpoint's own to absorb
}
}
}()
ts := time.Now().UTC()
for i := 0; i < contentionIterations; i++ {
asns := map[int]time.Time{
i % contendedASNCount: ts,
(i % contendedASNCount) + asnSecondBand: ts,
}
if err := db.GetOrCreateASNBatch(asns); err != nil {
cancel()
wg.Wait()
t.Fatalf("batch write failed under checkpoint contention: %v", err)
}
}
cancel()
wg.Wait()
}
// TestStatsRouteTimestampsAndCounts checks GetStatsContext reports the correct
// route counts and the oldest/newest last_updated across both route tables. The
// old union-scan query read the aggregate result into *time.Time, which the
// driver could not parse, so it logged a warning every call and left both
// timestamps nil; this asserts they are populated from the right rows.
func TestStatsRouteTimestampsAndCounts(t *testing.T) {
cfg := &config.Config{StateDir: t.TempDir()}
db, err := New(cfg, logger.New())
if err != nil {
t.Fatalf("failed to create database: %v", err)
}
defer func() { _ = db.Close() }()
ctx := context.Background()
// Empty database: no routes, so both timestamps are nil and no error.
empty, err := db.GetStatsContext(ctx)
if err != nil {
t.Fatalf("GetStatsContext on empty database: %v", err)
}
if empty.OldestRoute != nil || empty.NewestRoute != nil {
t.Fatalf("empty database timestamps = (%v, %v), want (nil, nil)",
empty.OldestRoute, empty.NewestRoute)
}
if empty.LiveRoutes != 0 {
t.Fatalf("empty database LiveRoutes = %d, want 0", empty.LiveRoutes)
}
base := time.Date(2026, 1, 2, 3, 4, 5, 0, time.UTC)
oldest := base
middle := base.Add(time.Minute)
newest := base.Add(2 * time.Minute)
mkV4 := func(prefix string, asn int, ts time.Time) *LiveRoute {
start, end, rerr := CalculateIPv4Range(prefix)
if rerr != nil {
t.Fatalf("CalculateIPv4Range(%s): %v", prefix, rerr)
}
return &LiveRoute{
ID: uuid.New(),
Prefix: prefix,
MaskLength: 24,
IPVersion: ipVersionV4,
OriginASN: asn,
PeerIP: "192.0.2.1",
ASPath: []int{asn},
NextHop: "192.0.2.254",
LastUpdated: ts,
V4IPStart: &start,
V4IPEnd: &end,
}
}
// Two IPv4 routes (one oldest, one middle) and one IPv6 route (newest).
routes := []*LiveRoute{
mkV4("198.51.100.0/24", 64500, middle),
mkV4("203.0.113.0/24", 64501, oldest),
{
ID: uuid.New(),
Prefix: "2001:db8::/32",
MaskLength: 32,
IPVersion: ipVersionV6,
OriginASN: 64502,
PeerIP: "2001:db8::1",
ASPath: []int{64502},
NextHop: "2001:db8::ffff",
LastUpdated: newest,
},
}
for _, route := range routes {
if err := db.UpsertLiveRoute(route); err != nil {
t.Fatalf("UpsertLiveRoute(%s): %v", route.Prefix, err)
}
}
stats, err := db.GetStatsContext(ctx)
if err != nil {
t.Fatalf("GetStatsContext: %v", err)
}
if stats.IPv4Routes != 2 {
t.Errorf("IPv4Routes = %d, want 2", stats.IPv4Routes)
}
if stats.IPv6Routes != 1 {
t.Errorf("IPv6Routes = %d, want 1", stats.IPv6Routes)
}
if stats.LiveRoutes != 3 {
t.Errorf("LiveRoutes = %d, want 3", stats.LiveRoutes)
}
if stats.OldestRoute == nil || !stats.OldestRoute.Equal(oldest) {
t.Errorf("OldestRoute = %v, want %v", stats.OldestRoute, oldest)
}
if stats.NewestRoute == nil || !stats.NewestRoute.Equal(newest) {
t.Errorf("NewestRoute = %v, want %v", stats.NewestRoute, newest)
}
}
func BenchmarkIPToUint32(b *testing.B) {
ip := net.ParseIP("192.168.1.1")
b.ResetTimer()
for i := 0; i < b.N; i++ {
_ = ipToUint32(ip)
}
}
func BenchmarkCalculateIPv4Range(b *testing.B) {
cidr := "192.168.0.0/16"
b.ResetTimer()
for i := 0; i < b.N; i++ {
_, _, _ = CalculateIPv4Range(cidr)
}
}