package aggregate import ( "testing" "time" "github.com/netsynth/netsynth/classify" ) // sendPackets sends a slice of ClassifiedPackets to the events channel. func sendPackets(events chan<- classify.ClassifiedPacket, packets []classify.ClassifiedPacket) { for _, p := range packets { events <- p } } func TestAggregateEmitsSnapshot(t *testing.T) { done := make(chan struct{}) events := make(chan classify.ClassifiedPacket, 20) defer close(done) // Use short window to trigger ticks quickly. out := Aggregate(done, events, 20, nil) // Send 5 ICMP packets. for i := 0; i < 5; i++ { events <- classify.ClassifiedPacket{Class: classify.ClassICMP} } // Wait for at least one snapshot from a tick. select { case snap := <-out: if snap.TotalPackets < 5 { // Packets may span windows; total across first snapshot should be <= 5. // Just verify snapshot is emitted and counts are non-negative. } _ = snap case <-time.After(500 * time.Millisecond): t.Fatal("timeout: no snapshot emitted within 500ms") } } func TestAggregateMultipleClasses(t *testing.T) { done := make(chan struct{}) events := make(chan classify.ClassifiedPacket, 20) out := Aggregate(done, events, 20, nil) events <- classify.ClassifiedPacket{Class: classify.ClassICMP} events <- classify.ClassifiedPacket{Class: classify.ClassICMP} events <- classify.ClassifiedPacket{Class: classify.ClassDNS} events <- classify.ClassifiedPacket{Class: classify.ClassHTTPS} events <- classify.ClassifiedPacket{Class: classify.ClassHTTPS} events <- classify.ClassifiedPacket{Class: classify.ClassHTTPS} // Close done to get the final flush. // But first drain any tick-based snapshots. time.Sleep(10 * time.Millisecond) close(done) // Accumulate all snapshots to verify totals. totals := make(map[classify.TrafficClass]int64) for snap := range out { for class, count := range snap.Counts { totals[class] += count } } if totals[classify.ClassICMP] != 2 { t.Errorf("expected ICMP=2, got %d", totals[classify.ClassICMP]) } if totals[classify.ClassDNS] != 1 { t.Errorf("expected DNS=1, got %d", totals[classify.ClassDNS]) } if totals[classify.ClassHTTPS] != 3 { t.Errorf("expected HTTPS=3, got %d", totals[classify.ClassHTTPS]) } } func TestAggregateDoneFlushesPartial(t *testing.T) { done := make(chan struct{}) // Use very large window so no tick fires. events := make(chan classify.ClassifiedPacket, 20) out := Aggregate(done, events, 100_000, nil) events <- classify.ClassifiedPacket{Class: classify.ClassSSH} events <- classify.ClassifiedPacket{Class: classify.ClassSSH} // Give goroutine time to process the two events. time.Sleep(20 * time.Millisecond) // Close done to trigger flush. close(done) var snap classify.WindowSnapshot select { case snap = <-out: case <-time.After(500 * time.Millisecond): t.Fatal("timeout: no snapshot on done-flush") } if snap.Counts[classify.ClassSSH] != 2 { t.Errorf("expected SSH=2 in flush, got %d", snap.Counts[classify.ClassSSH]) } if snap.TotalPackets != 2 { t.Errorf("expected TotalPackets=2, got %d", snap.TotalPackets) } } func TestAggregateEmptyWindow(t *testing.T) { done := make(chan struct{}) events := make(chan classify.ClassifiedPacket, 20) out := Aggregate(done, events, 20, nil) // Don't send any packets, just wait for a tick. select { case snap := <-out: if snap.TotalPackets != 0 { t.Errorf("expected TotalPackets=0 for empty window, got %d", snap.TotalPackets) } if len(snap.Counts) != 0 { t.Errorf("expected empty Counts for empty window, got %v", snap.Counts) } case <-time.After(500 * time.Millisecond): t.Fatal("timeout: no snapshot for empty window") } close(done) } func TestAggregateWindowIndex(t *testing.T) { done := make(chan struct{}) events := make(chan classify.ClassifiedPacket, 5) out := Aggregate(done, events, 20, nil) // Collect 3 snapshots. var indices []int timeout := time.After(2 * time.Second) for len(indices) < 3 { select { case snap := <-out: indices = append(indices, snap.WindowIndex) case <-timeout: t.Fatalf("timeout: only collected %d snapshots", len(indices)) } } close(done) // WindowIndex should increment: 0, 1, 2. for i, idx := range indices { if idx != i { t.Errorf("expected WindowIndex[%d]=%d, got %d", i, i, idx) } } } // --- AggregatePcap tests --- // baseTime is a fixed reference time for pcap aggregation tests. var baseTime = time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC) // makePcapChan creates a buffered channel, sends the given packets, closes it, and returns it. func makePcapChan(pkts []classify.ClassifiedPacket) <-chan classify.ClassifiedPacket { ch := make(chan classify.ClassifiedPacket, len(pkts)+1) for _, p := range pkts { ch <- p } close(ch) return ch } // TestAggregatePcapBasic verifies 3 packets in window 0 produce 1 snapshot with TotalPackets=3. func TestAggregatePcapBasic(t *testing.T) { pkts := []classify.ClassifiedPacket{ {Class: classify.ClassICMP, Timestamp: baseTime}, {Class: classify.ClassICMP, Timestamp: baseTime.Add(100 * time.Millisecond)}, {Class: classify.ClassICMP, Timestamp: baseTime.Add(200 * time.Millisecond)}, } snaps := AggregatePcap(makePcapChan(pkts), 500, nil) if len(snaps) != 1 { t.Fatalf("expected 1 snapshot, got %d", len(snaps)) } if snaps[0].TotalPackets != 3 { t.Errorf("expected TotalPackets=3, got %d", snaps[0].TotalPackets) } } // TestAggregatePcapMultipleWindows verifies 3 packets spanning 3 windows each produce 1 packet. func TestAggregatePcapMultipleWindows(t *testing.T) { pkts := []classify.ClassifiedPacket{ {Class: classify.ClassDNS, Timestamp: baseTime}, {Class: classify.ClassDNS, Timestamp: baseTime.Add(600 * time.Millisecond)}, {Class: classify.ClassDNS, Timestamp: baseTime.Add(1200 * time.Millisecond)}, } snaps := AggregatePcap(makePcapChan(pkts), 500, nil) if len(snaps) != 3 { t.Fatalf("expected 3 snapshots, got %d", len(snaps)) } for i, s := range snaps { if s.TotalPackets != 1 { t.Errorf("snapshot[%d] expected TotalPackets=1, got %d", i, s.TotalPackets) } } } // TestAggregatePcapGaps verifies that gap windows produce empty snapshots (D-02: gaps are silent). func TestAggregatePcapGaps(t *testing.T) { // Packets at T+0ms and T+1500ms (skipping windows 1 and 2) pkts := []classify.ClassifiedPacket{ {Class: classify.ClassICMP, Timestamp: baseTime}, {Class: classify.ClassICMP, Timestamp: baseTime.Add(1500 * time.Millisecond)}, } snaps := AggregatePcap(makePcapChan(pkts), 500, nil) // Expected: window 0 (1 pkt), window 1 (0 pkts), window 2 (0 pkts), window 3 (1 pkt) = 4 snapshots if len(snaps) != 4 { t.Fatalf("expected 4 snapshots (with gaps), got %d", len(snaps)) } if snaps[0].TotalPackets != 1 { t.Errorf("snapshot[0] expected TotalPackets=1, got %d", snaps[0].TotalPackets) } if snaps[1].TotalPackets != 0 { t.Errorf("snapshot[1] expected TotalPackets=0 (gap), got %d", snaps[1].TotalPackets) } if snaps[2].TotalPackets != 0 { t.Errorf("snapshot[2] expected TotalPackets=0 (gap), got %d", snaps[2].TotalPackets) } if snaps[3].TotalPackets != 1 { t.Errorf("snapshot[3] expected TotalPackets=1, got %d", snaps[3].TotalPackets) } } // TestAggregatePcapEmpty verifies that empty channel returns empty (nil) slice. func TestAggregatePcapEmpty(t *testing.T) { snaps := AggregatePcap(makePcapChan(nil), 500, nil) if len(snaps) != 0 { t.Errorf("expected 0 snapshots for empty input, got %d", len(snaps)) } } // TestAggregatePcapWindowIndex verifies sequential WindowIndex values. func TestAggregatePcapWindowIndex(t *testing.T) { pkts := []classify.ClassifiedPacket{ {Class: classify.ClassHTTPS, Timestamp: baseTime}, {Class: classify.ClassHTTPS, Timestamp: baseTime.Add(600 * time.Millisecond)}, {Class: classify.ClassHTTPS, Timestamp: baseTime.Add(1200 * time.Millisecond)}, } snaps := AggregatePcap(makePcapChan(pkts), 500, nil) for i, s := range snaps { if s.WindowIndex != i { t.Errorf("snapshot[%d].WindowIndex = %d; want %d", i, s.WindowIndex, i) } } } // TestAggregatePcapClassCounts verifies per-class counts in a shared window. func TestAggregatePcapClassCounts(t *testing.T) { pkts := []classify.ClassifiedPacket{ {Class: classify.ClassICMP, Timestamp: baseTime}, {Class: classify.ClassDNS, Timestamp: baseTime.Add(50 * time.Millisecond)}, {Class: classify.ClassICMP, Timestamp: baseTime.Add(100 * time.Millisecond)}, {Class: classify.ClassHTTPS, Timestamp: baseTime.Add(150 * time.Millisecond)}, } snaps := AggregatePcap(makePcapChan(pkts), 500, nil) if len(snaps) != 1 { t.Fatalf("expected 1 snapshot, got %d", len(snaps)) } s := snaps[0] if s.Counts[classify.ClassICMP] != 2 { t.Errorf("expected ICMP=2, got %d", s.Counts[classify.ClassICMP]) } if s.Counts[classify.ClassDNS] != 1 { t.Errorf("expected DNS=1, got %d", s.Counts[classify.ClassDNS]) } if s.Counts[classify.ClassHTTPS] != 1 { t.Errorf("expected HTTPS=1, got %d", s.Counts[classify.ClassHTTPS]) } if s.TotalPackets != 4 { t.Errorf("expected TotalPackets=4, got %d", s.TotalPackets) } } // TestAggregatePcapOnSnapshot verifies that the onSnapshot callback fires for each snapshot. func TestAggregatePcapOnSnapshot(t *testing.T) { pkts := []classify.ClassifiedPacket{ {Class: classify.ClassICMP, Timestamp: baseTime}, {Class: classify.ClassDNS, Timestamp: baseTime.Add(600 * time.Millisecond)}, } var called int var calledIndices []int onSnapshot := func(s classify.WindowSnapshot) { called++ calledIndices = append(calledIndices, s.WindowIndex) } snaps := AggregatePcap(makePcapChan(pkts), 500, onSnapshot) if called != len(snaps) { t.Errorf("onSnapshot called %d times; want %d (once per snapshot)", called, len(snaps)) } for i, idx := range calledIndices { if idx != i { t.Errorf("onSnapshot call[%d] had WindowIndex=%d; want %d", i, idx, i) } } }