package database import ( "context" "database/sql" "net" "sync" "testing" "time" "git.eeqj.de/sneak/routewatch/internal/config" "git.eeqj.de/sneak/routewatch/internal/logger" ) // 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() } 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) } }