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) } }