Files
yoloyolo/aggregate/window_test.go
T

307 lines
9.7 KiB
Go
Raw Normal View History

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