diff --git a/README.md b/README.md index 85971fe..ca079dd 100644 --- a/README.md +++ b/README.md @@ -9,7 +9,6 @@ [![License](https://img.shields.io/badge/license-MIT-blue.svg)](https://github.com/v-byte-cpu/sx/blob/master/LICENSE) [![Build Status](https://github.com/v-byte-cpu/sx/actions/workflows/ci.yml/badge.svg)](https://github.com/v-byte-cpu/sx/actions/workflows/ci.yml) -[![GoReportCard Status](https://goreportcard.com/badge/github.com/v-byte-cpu/sx)](https://goreportcard.com/report/github.com/v-byte-cpu/sx) ![Platform](https://img.shields.io/badge/platform-linux%2FmacOS%2Fdocker-blue) diff --git a/command/config.go b/command/config.go index e60d81d..83c7d34 100644 --- a/command/config.go +++ b/command/config.go @@ -462,7 +462,7 @@ func (o *genericScanCmdOpts) getLogger(name string, w io.Writer) (logger log.Log return } -func (o *genericScanCmdOpts) newScanEngine(ctx context.Context, scanner scan.Scanner) *scan.GenericEngine { +func (o *genericScanCmdOpts) newScanEngine(ctx context.Context, scanner scan.Scanner) *scan.ScanEngine { if o.rateCount > 0 { scanner = scan.NewRateLimitScanner(scanner, ratelimit.New(o.rateCount, ratelimit.Per(o.rateWindow))) diff --git a/pkg/packet/mock_sender_test.go b/pkg/packet/mock_sender_test.go index 227c2d0..a106d42 100644 --- a/pkg/packet/mock_sender_test.go +++ b/pkg/packet/mock_sender_test.go @@ -41,10 +41,10 @@ func (m *MockSender) EXPECT() *MockSenderMockRecorder { } // SendPackets mocks base method. -func (m *MockSender) SendPackets(ctx context.Context, in <-chan *BufferData) (<-chan any, <-chan error) { +func (m *MockSender) SendPackets(ctx context.Context, in <-chan *BufferData) (<-chan struct{}, <-chan error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "SendPackets", ctx, in) - ret0, _ := ret[0].(<-chan any) + ret0, _ := ret[0].(<-chan struct{}) ret1, _ := ret[1].(<-chan error) return ret0, ret1 } diff --git a/pkg/packet/receiver_test.go b/pkg/packet/receiver_test.go index c77ad8b..ab4a7f4 100644 --- a/pkg/packet/receiver_test.go +++ b/pkg/packet/receiver_test.go @@ -9,7 +9,6 @@ import ( "testing" "github.com/google/gopacket" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" ) @@ -65,8 +64,8 @@ func TestReceivePacketsWithUnrecoverableError(t *testing.T) { r := NewReceiver(sr, p) out := r.ReceivePackets(context.Background()) - result := chanToSlice(t, chanErrToGeneric(out), 0) - assert.Empty(t, result, "error slice is not empty") + result := chanToSlice(t, out, 0) + require.Empty(t, result, "error slice is not empty") }) } } @@ -91,8 +90,8 @@ func TestReceivePacketsOnePacket(t *testing.T) { r := NewReceiver(sr, p) out := r.ReceivePackets(context.Background()) - result := chanToSlice(t, chanErrToGeneric(out), 0) - assert.Empty(t, result, "error slice is not empty") + result := chanToSlice(t, out, 0) + require.Empty(t, result, "error slice is not empty") } func TestReceivePacketsOnePacketWithProcessError(t *testing.T) { @@ -114,9 +113,9 @@ func TestReceivePacketsOnePacketWithProcessError(t *testing.T) { r := NewReceiver(sr, p) out := r.ReceivePackets(context.Background()) - result := chanToSlice(t, chanErrToGeneric(out), 1) - assert.Len(t, result, 1, "error slice is invalid") - require.Error(t, result[0].(error)) + result := chanToSlice(t, out, 1) + require.Len(t, result, 1, "error slice is invalid") + require.Error(t, result[0]) } func TestReceivePacketsOnePacketWithRetryError(t *testing.T) { @@ -158,8 +157,8 @@ func TestReceivePacketsOnePacketWithRetryError(t *testing.T) { r := NewReceiver(sr, p) out := r.ReceivePackets(context.Background()) - result := chanToSlice(t, chanErrToGeneric(out), 0) - assert.Empty(t, result, "error slice is not empty") + result := chanToSlice(t, out, 0) + require.Empty(t, result, "error slice is not empty") }) } } @@ -185,9 +184,9 @@ func TestReceivePacketsOnePacketWithUnknownError(t *testing.T) { r := NewReceiver(sr, p) out := r.ReceivePackets(context.Background()) - result := chanToSlice(t, chanErrToGeneric(out), 1) - assert.Len(t, result, 1, "error slice length is invalid") - require.Error(t, result[0].(error)) + result := chanToSlice(t, out, 1) + require.Len(t, result, 1, "error slice length is invalid") + require.Error(t, result[0]) } func TestReceivePacketsOnePacketWithContextCancel(t *testing.T) { @@ -212,6 +211,6 @@ func TestReceivePacketsOnePacketWithContextCancel(t *testing.T) { r := NewReceiver(sr, p) out := r.ReceivePackets(ctx) - result := chanToSlice(t, chanErrToGeneric(out), 0) - assert.Empty(t, result, "error slice is not empty") + result := chanToSlice(t, out, 0) + require.Empty(t, result, "error slice is not empty") } diff --git a/pkg/packet/sender.go b/pkg/packet/sender.go index 2db79a2..f356fa6 100644 --- a/pkg/packet/sender.go +++ b/pkg/packet/sender.go @@ -13,8 +13,10 @@ type BufferData struct { Err error } +// Sender writes serialized packets. type Sender interface { - SendPackets(ctx context.Context, in <-chan *BufferData) (done <-chan interface{}, errc <-chan error) + // SendPackets writes every packet from in until completion or context cancellation. + SendPackets(ctx context.Context, in <-chan *BufferData) (done <-chan struct{}, errc <-chan error) } type Writer interface { @@ -29,8 +31,8 @@ type sender struct { w Writer } -func (s *sender) SendPackets(ctx context.Context, in <-chan *BufferData) (<-chan interface{}, <-chan error) { - done := make(chan interface{}) +func (s *sender) SendPackets(ctx context.Context, in <-chan *BufferData) (<-chan struct{}, <-chan error) { + done := make(chan struct{}) errc := make(chan error, 100) go func() { defer func() { @@ -45,18 +47,39 @@ func (s *sender) SendPackets(ctx context.Context, in <-chan *BufferData) (<-chan if !ok { return } - if pkt.Err != nil { - errc <- pkt.Err - continue - } - if err := s.w.WritePacketData(pkt.Buf.Bytes()); err != nil { - errc <- err - } - if err := FreeSerializeBuffer(pkt.Buf); err != nil { - errc <- err + if !s.sendPacket(ctx, errc, pkt) { + return } } } }() return done, errc } + +func (s *sender) sendPacket(ctx context.Context, errc chan<- error, pkt *BufferData) bool { + if pkt.Err != nil { + return sendError(ctx, errc, pkt.Err) + } + if err := s.w.WritePacketData(pkt.Buf.Bytes()); err != nil { + if !sendError(ctx, errc, err) { + _ = FreeSerializeBuffer(pkt.Buf) + return false + } + } + if err := FreeSerializeBuffer(pkt.Buf); err != nil { + return sendError(ctx, errc, err) + } + return true +} + +func sendError(ctx context.Context, out chan<- error, err error) bool { + if ctx.Err() != nil { + return false + } + select { + case <-ctx.Done(): + return false + case out <- err: + return true + } +} diff --git a/pkg/packet/sender_test.go b/pkg/packet/sender_test.go index 9c8e663..68d8999 100644 --- a/pkg/packet/sender_test.go +++ b/pkg/packet/sender_test.go @@ -7,11 +7,17 @@ import ( "time" "github.com/google/gopacket" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" ) +type typedDoneSender interface { + // SendPackets exposes the expected typed completion contract. + SendPackets(ctx context.Context, in <-chan *BufferData) (done <-chan struct{}, errc <-chan error) +} + +var _ typedDoneSender = (Sender)(nil) + func TestSenderWithEmptyChannel(t *testing.T) { t.Parallel() in := make(chan *BufferData) @@ -23,10 +29,10 @@ func TestSenderWithEmptyChannel(t *testing.T) { done, errc := s.SendPackets(context.Background(), in) - result := chanToSlice(t, chanErrToGeneric(errc), 0) - assert.Empty(t, result, "error slice is not empty") - result = chanToSlice(t, done, 0) - assert.Empty(t, result, "error slice is not empty") + result := chanToSlice(t, errc, 0) + require.Empty(t, result, "error slice is not empty") + doneResult := chanToSlice(t, done, 0) + require.Empty(t, doneResult, "done channel is not empty") } func TestSenderWithOnePacket(t *testing.T) { @@ -49,10 +55,10 @@ func TestSenderWithOnePacket(t *testing.T) { done, errc := s.SendPackets(context.Background(), in) - result := chanToSlice(t, chanErrToGeneric(errc), 0) - assert.Empty(t, result, "error slice is not empty") - result = chanToSlice(t, done, 0) - assert.Empty(t, result, "error slice is not empty") + result := chanToSlice(t, errc, 0) + require.Empty(t, result, "error slice is not empty") + doneResult := chanToSlice(t, done, 0) + require.Empty(t, doneResult, "done channel is not empty") } func TestSenderWithTwoPackets(t *testing.T) { @@ -88,10 +94,10 @@ func TestSenderWithTwoPackets(t *testing.T) { done, errc := s.SendPackets(context.Background(), in) - result := chanToSlice(t, chanErrToGeneric(errc), 0) - assert.Empty(t, result, "error slice is not empty") - result = chanToSlice(t, done, 0) - assert.Empty(t, result, "error slice is not empty") + result := chanToSlice(t, errc, 0) + require.Empty(t, result, "error slice is not empty") + doneResult := chanToSlice(t, done, 0) + require.Empty(t, doneResult, "done channel is not empty") } func TestSenderWithInvalidPacketReturnsError(t *testing.T) { @@ -106,12 +112,12 @@ func TestSenderWithInvalidPacketReturnsError(t *testing.T) { done, errc := s.SendPackets(context.Background(), in) - result := chanToSlice(t, chanErrToGeneric(errc), 1) - assert.Len(t, result, 1, "error slice size is invalid") - require.Error(t, result[0].(error)) + result := chanToSlice(t, errc, 1) + require.Len(t, result, 1, "error slice size is invalid") + require.Error(t, result[0]) - result = chanToSlice(t, done, 0) - assert.Empty(t, result, "error slice is not empty") + doneResult := chanToSlice(t, done, 0) + require.Empty(t, doneResult, "done channel is not empty") } func TestSenderWithWriteErrorReturnsError(t *testing.T) { @@ -132,12 +138,12 @@ func TestSenderWithWriteErrorReturnsError(t *testing.T) { done, errc := s.SendPackets(context.Background(), in) - result := chanToSlice(t, chanErrToGeneric(errc), 1) - assert.Len(t, result, 1, "error slice size is invalid") - require.Error(t, result[0].(error)) + result := chanToSlice(t, errc, 1) + require.Len(t, result, 1, "error slice size is invalid") + require.Error(t, result[0]) - result = chanToSlice(t, done, 0) - assert.Empty(t, result, "error slice is not empty") + doneResult := chanToSlice(t, done, 0) + require.Empty(t, doneResult, "done channel is not empty") } func TestSenderWithTimeout(t *testing.T) { @@ -157,5 +163,33 @@ func TestSenderWithTimeout(t *testing.T) { require.FailNow(t, "exit timeout") } result := chanToSlice(t, done, 0) - assert.Empty(t, result, "error slice is not empty") + require.Empty(t, result, "error slice is not empty") +} + +func TestSenderCancellationWithFullErrorChannel(t *testing.T) { + t.Parallel() + + const errorChannelCapacity = 100 + in := make(chan *BufferData, errorChannelCapacity+1) + for range errorChannelCapacity + 1 { + in <- &BufferData{Err: errors.New("invalid data")} + } + close(in) + + ctrl := gomock.NewController(t) + w := NewMockWriter(ctrl) + s := NewSender(w) + ctx, cancel := context.WithCancel(context.Background()) + done, errc := s.SendPackets(ctx, in) + + require.Eventually(t, func() bool { + return len(errc) == cap(errc) + }, time.Second, time.Millisecond) + cancel() + + select { + case <-done: + case <-time.After(time.Second): + require.FailNow(t, "sender did not stop after context cancellation") + } } diff --git a/pkg/packet/utils_test.go b/pkg/packet/utils_test.go index 14a2f60..c6bbf56 100644 --- a/pkg/packet/utils_test.go +++ b/pkg/packet/utils_test.go @@ -9,9 +9,9 @@ import ( const waitTimeout = 3 * time.Second -func chanToSlice(t *testing.T, in <-chan interface{}, expectedLen int) []interface{} { +func chanToSlice[T any](t *testing.T, in <-chan T, expectedLen int) []T { t.Helper() - result := []interface{}{} + result := []T{} loop: for { select { @@ -24,19 +24,8 @@ loop: } result = append(result, data) case <-time.After(waitTimeout): - t.Fatal("read timeout") + require.FailNow(t, "read timeout") } } return result } - -func chanErrToGeneric(in <-chan error) <-chan interface{} { - out := make(chan interface{}, cap(in)) - go func() { - defer close(out) - for i := range in { - out <- i - } - }() - return out -} diff --git a/pkg/scan/channels.go b/pkg/scan/channels.go new file mode 100644 index 0000000..c1ee6ac --- /dev/null +++ b/pkg/scan/channels.go @@ -0,0 +1,45 @@ +package scan + +import ( + "context" + "sync" +) + +func sendContext[T any](ctx context.Context, out chan<- T, value T) bool { + if ctx.Err() != nil { + return false + } + select { + case <-ctx.Done(): + return false + case out <- value: + return true + } +} + +func mergeChannels[T any](ctx context.Context, capacity int, channels ...<-chan T) <-chan T { + out := make(chan T, capacity) + var wg sync.WaitGroup + wg.Add(len(channels)) + + for _, channel := range channels { + go func() { + defer wg.Done() + for { + select { + case <-ctx.Done(): + return + case value, ok := <-channel: + if !ok || !sendContext(ctx, out, value) { + return + } + } + } + }() + } + go func() { + wg.Wait() + close(out) + }() + return out +} diff --git a/pkg/scan/engine.go b/pkg/scan/engine.go index 20cdceb..4998452 100644 --- a/pkg/scan/engine.go +++ b/pkg/scan/engine.go @@ -24,8 +24,10 @@ type Range struct { Ports []*PortRange } +// Engine runs a scan and reports its completion and errors. type Engine interface { - Start(ctx context.Context, r *Range) (done <-chan interface{}, errc <-chan error) + // Start runs a scan for r until completion or context cancellation. + Start(ctx context.Context, r *Range) (done <-chan struct{}, errc <-chan error) } type Resulter interface { @@ -80,41 +82,12 @@ func NewPacketEngine(ps PacketSource, s packet.Sender, r packet.Receiver) *Packe return &PacketEngine{src: ps, snd: s, rcv: r} } -func (e *PacketEngine) Start(ctx context.Context, r *Range) (<-chan interface{}, <-chan error) { +// Start runs a packet scan for r until completion or context cancellation. +func (e *PacketEngine) Start(ctx context.Context, r *Range) (<-chan struct{}, <-chan error) { packets := e.src.Packets(ctx, r) done, errc1 := e.snd.SendPackets(ctx, packets) errc2 := e.rcv.ReceivePackets(ctx) - return done, mergeErrChan(ctx, errc1, errc2) -} - -// generics would be helpful :) -func mergeErrChan(ctx context.Context, channels ...<-chan error) <-chan error { - var wg sync.WaitGroup - wg.Add(len(channels)) - - out := make(chan error, 100) - multiplex := func(c <-chan error) { - defer wg.Done() - for { - select { - case <-ctx.Done(): - return - case e, ok := <-c: - if !ok { - return - } - writeError(ctx, out, e) - } - } - } - for _, c := range channels { - go multiplex(c) - } - go func() { - wg.Wait() - close(out) - }() - return out + return done, mergeChannels(ctx, 100, errc1, errc2) } type PacketMethod interface { @@ -153,27 +126,31 @@ func (s *rateLimitScanner) Scan(ctx context.Context, r *Request) (Result, error) return s.Scanner.Scan(ctx, r) } -type GenericEngine struct { +// ScanEngine runs scanner workers over generated requests. +type ScanEngine struct { reqgen RequestGenerator scanner Scanner results ResultChan workerCount int } -// Assert that GenericEngine conforms to the scan.EngineResulter interface -var _ EngineResulter = (*GenericEngine)(nil) +// Assert that ScanEngine conforms to the scan.EngineResulter interface. +var _ EngineResulter = (*ScanEngine)(nil) -type GenericEngineOption func(s *GenericEngine) +// ScanEngineOption configures a ScanEngine. +type ScanEngineOption func(s *ScanEngine) -func WithScanWorkerCount(workerCount int) GenericEngineOption { - return func(s *GenericEngine) { +// WithScanWorkerCount sets the number of concurrent scanner workers. +func WithScanWorkerCount(workerCount int) ScanEngineOption { + return func(s *ScanEngine) { s.workerCount = workerCount } } +// NewScanEngine creates a ScanEngine from its request, scanner, and result dependencies. func NewScanEngine(reqgen RequestGenerator, - scanner Scanner, results ResultChan, opts ...GenericEngineOption) *GenericEngine { - s := &GenericEngine{ + scanner Scanner, results ResultChan, opts ...ScanEngineOption) *ScanEngine { + s := &ScanEngine{ reqgen: reqgen, scanner: scanner, results: results, @@ -185,12 +162,14 @@ func NewScanEngine(reqgen RequestGenerator, return s } -func (e *GenericEngine) Results() <-chan Result { +// Results returns scan results until the result context is canceled. +func (e *ScanEngine) Results() <-chan Result { return e.results.Chan() } -func (e *GenericEngine) Start(ctx context.Context, r *Range) (<-chan interface{}, <-chan error) { - done := make(chan interface{}) +// Start runs a scan for r until completion or context cancellation. +func (e *ScanEngine) Start(ctx context.Context, r *Range) (<-chan struct{}, <-chan error) { + done := make(chan struct{}) errc := make(chan error, 100) requests, err := e.reqgen.GenerateRequests(ctx, r) if err != nil { @@ -212,7 +191,7 @@ func (e *GenericEngine) Start(ctx context.Context, r *Range) (<-chan interface{} return done, errc } -func (e *GenericEngine) worker(ctx context.Context, wg *sync.WaitGroup, +func (e *ScanEngine) worker(ctx context.Context, wg *sync.WaitGroup, requests <-chan *Request, errc chan<- error) { defer wg.Done() for { @@ -224,12 +203,16 @@ func (e *GenericEngine) worker(ctx context.Context, wg *sync.WaitGroup, return } if r.Err != nil { - writeError(ctx, errc, r.Err) + if !sendContext(ctx, errc, r.Err) { + return + } continue } result, err := e.scanner.Scan(ctx, r) if err != nil { - writeError(ctx, errc, err) + if !sendContext(ctx, errc, err) { + return + } continue } if result != nil { @@ -238,11 +221,3 @@ func (e *GenericEngine) worker(ctx context.Context, wg *sync.WaitGroup, } } } - -func writeError(ctx context.Context, out chan<- error, err error) { - select { - case <-ctx.Done(): - return - case out <- err: - } -} diff --git a/pkg/scan/engine_test.go b/pkg/scan/engine_test.go index 665d60d..41667d2 100644 --- a/pkg/scan/engine_test.go +++ b/pkg/scan/engine_test.go @@ -11,27 +11,34 @@ import ( "time" "github.com/google/gopacket" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/v-byte-cpu/sx/pkg/packet" "go.uber.org/mock/gomock" "go.uber.org/ratelimit" ) -func TestMergeErrChanEmptyChannels(t *testing.T) { +type typedDoneEngine interface { + // Start exposes the expected typed completion contract. + Start(ctx context.Context, r *Range) (done <-chan struct{}, errc <-chan error) +} + +var _ typedDoneEngine = (Engine)(nil) +var _ EngineResulter = (*ScanEngine)(nil) + +func TestMergeChannelsErrorsEmptyChannels(t *testing.T) { t.Parallel() c1 := make(chan error) close(c1) c2 := make(chan error) close(c2) - out := mergeErrChan(context.Background(), c1, c2) - result := chanToSlice(t, chanErrToGeneric(out), 0) + out := mergeChannels(context.Background(), 100, c1, c2) + result := chanToSlice(t, out, 0) - assert.Empty(t, result, "error slice is not empty") + require.Empty(t, result, "error slice is not empty") } -func TestMergeErrChanOneElementAndEmptyChannel(t *testing.T) { +func TestMergeChannelsErrorsOneElementAndEmptyChannel(t *testing.T) { t.Parallel() c1 := make(chan error, 1) c1 <- errors.New("test error") @@ -39,14 +46,14 @@ func TestMergeErrChanOneElementAndEmptyChannel(t *testing.T) { c2 := make(chan error) close(c2) - out := mergeErrChan(context.Background(), c1, c2) - result := chanToSlice(t, chanErrToGeneric(out), 1) + out := mergeChannels(context.Background(), 100, c1, c2) + result := chanToSlice(t, out, 1) - assert.Len(t, result, 1, "error slice size is invalid") - require.Error(t, result[0].(error)) + require.Len(t, result, 1, "error slice size is invalid") + require.Error(t, result[0]) } -func TestMergeErrChanTwoElements(t *testing.T) { +func TestMergeChannelsErrorsTwoElements(t *testing.T) { t.Parallel() c1 := make(chan error, 1) c1 <- errors.New("test error") @@ -55,15 +62,15 @@ func TestMergeErrChanTwoElements(t *testing.T) { c2 <- errors.New("test error") close(c2) - out := mergeErrChan(context.Background(), c1, c2) - result := chanToSlice(t, chanErrToGeneric(out), 2) + out := mergeChannels(context.Background(), 100, c1, c2) + result := chanToSlice(t, out, 2) - assert.Len(t, result, 2, "error slice size is invalid") - require.Error(t, result[0].(error)) - assert.Error(t, result[1].(error)) + require.Len(t, result, 2, "error slice size is invalid") + require.Error(t, result[0]) + require.Error(t, result[1]) } -func TestMergeErrChanContextExit(t *testing.T) { +func TestMergeChannelsErrorsContextExit(t *testing.T) { t.Parallel() c1 := make(chan error) defer close(c1) @@ -73,10 +80,10 @@ func TestMergeErrChanContextExit(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 1*time.Millisecond) defer cancel() - out := mergeErrChan(ctx, c1, c2) - result := chanToSlice(t, chanErrToGeneric(out), 0) + out := mergeChannels(ctx, 100, c1, c2) + result := chanToSlice(t, out, 0) - assert.Empty(t, result, "error slice is not empty") + require.Empty(t, result, "error slice is not empty") } func TestPacketEngineStartCollectsAllErrors(t *testing.T) { @@ -115,252 +122,206 @@ func TestPacketEngineStartCollectsAllErrors(t *testing.T) { }, }) - result := chanToSlice(t, chanErrToGeneric(out), 2) - assert.Len(t, result, 2, "error slice is invalid") - require.Error(t, result[0].(error)) - assert.Error(t, result[1].(error)) + result := chanToSlice(t, out, 2) + require.Len(t, result, 2, "error slice is invalid") + require.Error(t, result[0]) + require.Error(t, result[1]) } func TestPacketSourceReturnsError(t *testing.T) { t.Parallel() - done := make(chan interface{}) - - go func() { - defer close(done) - - ctrl := gomock.NewController(t) - reqgen := NewMockRequestGenerator(ctrl) - pktgen := NewMockPacketGenerator(ctrl) + ctrl := gomock.NewController(t) + reqgen := NewMockRequestGenerator(ctrl) + pktgen := NewMockPacketGenerator(ctrl) - scanRange := &Range{ - SrcIP: net.IPv4(192, 168, 0, 1), - SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - Ports: []*PortRange{ - { - StartPort: 22, - EndPort: 22, - }, + scanRange := &Range{ + SrcIP: net.IPv4(192, 168, 0, 1), + SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, + Ports: []*PortRange{ + { + StartPort: 22, + EndPort: 22, }, - } + }, + } - reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), scanRange). - Return(nil, errors.New("generate error")) + reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), scanRange). + Return(nil, errors.New("generate error")) - ps := NewPacketSource(reqgen, pktgen) - out := ps.Packets(context.Background(), scanRange) - result := <-out - assert.Error(t, result.Err) - }() - waitDone(t, done) + ps := NewPacketSource(reqgen, pktgen) + results := chanToSlice(t, ps.Packets(context.Background(), scanRange), 1) + require.Error(t, results[0].Err) } func TestPacketSourceReturnsData(t *testing.T) { t.Parallel() - done := make(chan interface{}) - - go func() { - defer close(done) - - ctrl := gomock.NewController(t) - reqgen := NewMockRequestGenerator(ctrl) - pktgen := NewMockPacketGenerator(ctrl) + ctrl := gomock.NewController(t) + reqgen := NewMockRequestGenerator(ctrl) + pktgen := NewMockPacketGenerator(ctrl) - scanRange := &Range{ - SrcIP: net.IPv4(192, 168, 0, 1), - SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, - Ports: []*PortRange{ - { - StartPort: 22, - EndPort: 22, - }, + scanRange := &Range{ + SrcIP: net.IPv4(192, 168, 0, 1), + SrcMAC: net.HardwareAddr{0x1, 0x2, 0x3, 0x4, 0x5, 0x6}, + Ports: []*PortRange{ + { + StartPort: 22, + EndPort: 22, }, - } - requests := make(chan *Request) - close(requests) - reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), scanRange). - Return(requests, nil) - - data := &packet.BufferData{Buf: gopacket.NewSerializeBuffer()} - dataCh := make(chan *packet.BufferData, 1) - dataCh <- data - close(dataCh) - pktgen.EXPECT().Packets(gomock.Not(gomock.Nil()), requests).Return(dataCh) - - ps := NewPacketSource(reqgen, pktgen) - out := ps.Packets(context.Background(), scanRange) - result := <-out - assert.NoError(t, result.Err) - assert.Equal(t, data.Buf, result.Buf) - }() - waitDone(t, done) + }, + } + requests := make(chan *Request) + close(requests) + reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), scanRange). + Return(requests, nil) + + data := &packet.BufferData{Buf: gopacket.NewSerializeBuffer()} + dataCh := make(chan *packet.BufferData, 1) + dataCh <- data + close(dataCh) + pktgen.EXPECT().Packets(gomock.Not(gomock.Nil()), requests).Return(dataCh) + + ps := NewPacketSource(reqgen, pktgen) + results := chanToSlice(t, ps.Packets(context.Background(), scanRange), 1) + require.NoError(t, results[0].Err) + require.Equal(t, data.Buf, results[0].Buf) } func TestRateLimitScanner(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - ctrl := gomock.NewController(t) - scanner := NewMockScanner(ctrl) - - req1 := &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 22} - expectedResult := &mockScanResult{"id1"} - scanner.EXPECT().Scan(gomock.Not(gomock.Nil()), req1). - Return(expectedResult, nil).AnyTimes() - - rateScanner := NewRateLimitScanner(scanner, - ratelimit.New(2, ratelimit.Per(20*time.Millisecond))) - timer := time.After(10 * time.Millisecond) - count := 0 - loop: - for { - select { - case <-timer: - break loop - default: - result, err := rateScanner.Scan(context.Background(), req1) - assert.NoError(t, err) - assert.Equal(t, expectedResult, result) - count++ - } + ctrl := gomock.NewController(t) + scanner := NewMockScanner(ctrl) + + req1 := &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 22} + expectedResult := &mockScanResult{"id1"} + scanner.EXPECT().Scan(gomock.Not(gomock.Nil()), req1). + Return(expectedResult, nil).AnyTimes() + + rateScanner := NewRateLimitScanner(scanner, + ratelimit.New(2, ratelimit.Per(20*time.Millisecond))) + timer := time.After(10 * time.Millisecond) + count := 0 +loop: + for { + select { + case <-timer: + break loop + default: + result, err := rateScanner.Scan(context.Background(), req1) + require.NoError(t, err) + require.Equal(t, expectedResult, result) + count++ } - assert.LessOrEqual(t, count, 2) - }() - waitDone(t, done) + } + require.LessOrEqual(t, count, 2) } func TestScanEngineWithRequestGeneratorError(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - ctrl := gomock.NewController(t) - reqgen := NewMockRequestGenerator(ctrl) - scanner := NewMockScanner(ctrl) - ctx := context.Background() + ctrl := gomock.NewController(t) + reqgen := NewMockRequestGenerator(ctrl) + scanner := NewMockScanner(ctrl) + ctx := context.Background() - reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), &Range{}). - Return(nil, errors.New("generate error")) - engine := NewScanEngine(reqgen, scanner, NewResultChan(ctx, 10)) + reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), &Range{}). + Return(nil, errors.New("generate error")) + engine := NewScanEngine(reqgen, scanner, NewResultChan(ctx, 10)) - _, errc := engine.Start(ctx, &Range{}) - err := <-errc - assert.Error(t, err) - }() - waitDone(t, done) + _, errc := engine.Start(ctx, &Range{}) + err := <-errc + require.Error(t, err) } func TestScanEngineWithRequestError(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - ctrl := gomock.NewController(t) - reqgen := NewMockRequestGenerator(ctrl) - scanner := NewMockScanner(ctrl) - ctx := context.Background() - - requests := make(chan *Request, 1) - requests <- &Request{Err: errors.New("request error")} - close(requests) - reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), &Range{}). - Return(requests, nil) - engine := NewScanEngine(reqgen, scanner, NewResultChan(ctx, 10)) - - _, errc := engine.Start(ctx, &Range{}) - err := <-errc - assert.Error(t, err) - }() - waitDone(t, done) + ctrl := gomock.NewController(t) + reqgen := NewMockRequestGenerator(ctrl) + scanner := NewMockScanner(ctrl) + ctx := context.Background() + + requests := make(chan *Request, 1) + requests <- &Request{Err: errors.New("request error")} + close(requests) + reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), &Range{}). + Return(requests, nil) + engine := NewScanEngine(reqgen, scanner, NewResultChan(ctx, 10)) + + _, errc := engine.Start(ctx, &Range{}) + err := <-errc + require.Error(t, err) } func TestScanEngineWithScannerError(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - ctrl := gomock.NewController(t) - reqgen := NewMockRequestGenerator(ctrl) - scanner := NewMockScanner(ctrl) - ctx := context.Background() - - requests := make(chan *Request, 1) - req1 := &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 22} - requests <- req1 - close(requests) - reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), &Range{}). - Return(requests, nil) - scanner.EXPECT().Scan(gomock.Not(gomock.Nil()), req1).Return(nil, errors.New("scan error")) - engine := NewScanEngine(reqgen, scanner, NewResultChan(ctx, 10)) - - _, errc := engine.Start(ctx, &Range{}) - err := <-errc - assert.Error(t, err) - }() - waitDone(t, done) + ctrl := gomock.NewController(t) + reqgen := NewMockRequestGenerator(ctrl) + scanner := NewMockScanner(ctrl) + ctx := context.Background() + + requests := make(chan *Request, 1) + req1 := &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 22} + requests <- req1 + close(requests) + reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), &Range{}). + Return(requests, nil) + scanner.EXPECT().Scan(gomock.Not(gomock.Nil()), req1).Return(nil, errors.New("scan error")) + engine := NewScanEngine(reqgen, scanner, NewResultChan(ctx, 10)) + + _, errc := engine.Start(ctx, &Range{}) + err := <-errc + require.Error(t, err) } func TestScanEngineWithResults(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - ctrl := gomock.NewController(t) - reqgen := NewMockRequestGenerator(ctrl) - scanner := NewMockScanner(ctrl) - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() - - requests := make(chan *Request, 2) - req1 := &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 22} - req2 := &Request{DstIP: net.IPv4(192, 168, 0, 2), DstPort: 22} - requests <- req1 - requests <- req2 - close(requests) - reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), &Range{}). - Return(requests, nil) - - scanner.EXPECT().Scan(gomock.Not(gomock.Nil()), req1). - Return(&mockScanResult{"id1"}, nil) - scanner.EXPECT().Scan(gomock.Not(gomock.Nil()), req2). - Return(&mockScanResult{"id2"}, nil) - - resultCh := NewResultChan(ctx, 10) - engine := NewScanEngine(reqgen, scanner, resultCh, WithScanWorkerCount(10)) - - done, errc := engine.Start(ctx, &Range{}) - <-done - results := make([]Result, 2) - results[0] = <-resultCh.Chan() - results[1] = <-resultCh.Chan() - cancel() - assert.Empty(t, errc, "error channel is not empty") - result, ok := <-resultCh.Chan() - if ok { - assert.Fail(t, "result channel contains more elements than expected: ", result) - } + ctrl := gomock.NewController(t) + reqgen := NewMockRequestGenerator(ctrl) + scanner := NewMockScanner(ctrl) + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() - sort.Slice(results, func(i, j int) bool { - return results[i].ID() < results[j].ID() - }) - assert.Equal(t, []Result{ - &mockScanResult{"id1"}, - &mockScanResult{"id2"}, - }, results) - }() - waitDone(t, done) + requests := make(chan *Request, 2) + req1 := &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 22} + req2 := &Request{DstIP: net.IPv4(192, 168, 0, 2), DstPort: 22} + requests <- req1 + requests <- req2 + close(requests) + reqgen.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), &Range{}). + Return(requests, nil) + + scanner.EXPECT().Scan(gomock.Not(gomock.Nil()), req1). + Return(&mockScanResult{"id1"}, nil) + scanner.EXPECT().Scan(gomock.Not(gomock.Nil()), req2). + Return(&mockScanResult{"id2"}, nil) + + resultCh := NewResultChan(ctx, 10) + engine := NewScanEngine(reqgen, scanner, resultCh, WithScanWorkerCount(10)) + + done, errc := engine.Start(ctx, &Range{}) + <-done + results := make([]Result, 2) + results[0] = <-resultCh.Chan() + results[1] = <-resultCh.Chan() + cancel() + require.Empty(t, errc, "error channel is not empty") + result, ok := <-resultCh.Chan() + if ok { + require.Fail(t, "result channel contains more elements than expected: ", result) + } + + sort.Slice(results, func(i, j int) bool { + return results[i].ID() < results[j].ID() + }) + require.Equal(t, []Result{ + &mockScanResult{"id1"}, + &mockScanResult{"id2"}, + }, results) } type mockScanResult struct { diff --git a/pkg/scan/generator.go b/pkg/scan/generator.go index 20ceffd..a42dea2 100644 --- a/pkg/scan/generator.go +++ b/pkg/scan/generator.go @@ -4,7 +4,6 @@ package scan import ( "context" - "sync" "github.com/google/gopacket" "github.com/v-byte-cpu/sx/pkg/packet" @@ -38,28 +37,24 @@ func (g *packetGenerator) Packets(ctx context.Context, in <-chan *Request) <-cha if !ok { return } - if r.Err != nil { - writeBufToChan(ctx, out, &packet.BufferData{Err: r.Err}) - continue - } - buf := packet.NewSerializeBuffer() - if err := g.filler.Fill(buf, r); err != nil { - writeBufToChan(ctx, out, &packet.BufferData{Err: err}) - continue + if !g.sendPacket(ctx, out, r) { + return } - writeBufToChan(ctx, out, &packet.BufferData{Buf: buf}) } } }() return out } -func writeBufToChan(ctx context.Context, out chan *packet.BufferData, buf *packet.BufferData) { - select { - case <-ctx.Done(): - return - case out <- buf: +func (g *packetGenerator) sendPacket(ctx context.Context, out chan<- *packet.BufferData, r *Request) bool { + if r.Err != nil { + return sendContext(ctx, out, &packet.BufferData{Err: r.Err}) + } + buf := packet.NewSerializeBuffer() + if err := g.filler.Fill(buf, r); err != nil { + return sendContext(ctx, out, &packet.BufferData{Err: err}) } + return sendContext(ctx, out, &packet.BufferData{Buf: buf}) } func NewPacketMultiGenerator(filler PacketFiller, numWorkers int) PacketGenerator { @@ -77,39 +72,5 @@ func (g *packetMultiGenerator) Packets(ctx context.Context, in <-chan *Request) for i := 0; i < g.numWorkers; i++ { workers[i] = g.gen.Packets(ctx, in) } - return MergeBufferDataChan(ctx, workers...) -} - -// generics would be helpful :) -func MergeBufferDataChan(ctx context.Context, channels ...<-chan *packet.BufferData) <-chan *packet.BufferData { - var wg sync.WaitGroup - wg.Add(len(channels)) - - out := make(chan *packet.BufferData, len(channels)*100) - multiplex := func(c <-chan *packet.BufferData) { - defer wg.Done() - for { - select { - case <-ctx.Done(): - return - case e, ok := <-c: - if !ok { - return - } - select { - case <-ctx.Done(): - return - case out <- e: - } - } - } - } - for _, c := range channels { - go multiplex(c) - } - go func() { - wg.Wait() - close(out) - }() - return out + return mergeChannels(ctx, len(workers)*100, workers...) } diff --git a/pkg/scan/generator_test.go b/pkg/scan/generator_test.go index 37f9417..a14dbaf 100644 --- a/pkg/scan/generator_test.go +++ b/pkg/scan/generator_test.go @@ -8,23 +8,11 @@ import ( "testing" "time" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/v-byte-cpu/sx/pkg/packet" "go.uber.org/mock/gomock" ) -func chanBufferDataToGeneric(in <-chan *packet.BufferData) <-chan interface{} { - out := make(chan interface{}, cap(in)) - go func() { - defer close(out) - for i := range in { - out <- i - } - }() - return out -} - func TestGeneratorPacketsWithEmptyChannel(t *testing.T) { t.Parallel() in := make(chan *Request) @@ -35,8 +23,8 @@ func TestGeneratorPacketsWithEmptyChannel(t *testing.T) { g := NewPacketGenerator(f) out := g.Packets(context.Background(), in) - result := chanToSlice(t, chanBufferDataToGeneric(out), 0) - assert.Empty(t, result, "result is not empty") + result := chanToSlice(t, out, 0) + require.Empty(t, result, "result is not empty") } func TestMultiGeneratorPacketsWithEmptyChannel(t *testing.T) { @@ -49,8 +37,8 @@ func TestMultiGeneratorPacketsWithEmptyChannel(t *testing.T) { g := NewPacketMultiGenerator(f, runtime.NumCPU()) out := g.Packets(context.Background(), in) - result := chanToSlice(t, chanBufferDataToGeneric(out), 0) - assert.Empty(t, result, "result is not empty") + result := chanToSlice(t, out, 0) + require.Empty(t, result, "result is not empty") } func TestGeneratorPacketsWithOnePair(t *testing.T) { @@ -70,12 +58,12 @@ func TestGeneratorPacketsWithOnePair(t *testing.T) { g := NewPacketGenerator(f) out := g.Packets(context.Background(), in) - results := chanToSlice(t, chanBufferDataToGeneric(out), 1) + results := chanToSlice(t, out, 1) - assert.Len(t, results, 1, "result size is invalid") - result := results[0].(*packet.BufferData) + require.Len(t, results, 1, "result size is invalid") + result := results[0] require.NoError(t, result.Err) - assert.NotNil(t, result.Buf) + require.NotNil(t, result.Buf) } func TestMultiGeneratorPacketsWithOnePair(t *testing.T) { @@ -95,12 +83,12 @@ func TestMultiGeneratorPacketsWithOnePair(t *testing.T) { g := NewPacketMultiGenerator(f, runtime.NumCPU()) out := g.Packets(context.Background(), in) - results := chanToSlice(t, chanBufferDataToGeneric(out), 1) + results := chanToSlice(t, out, 1) - assert.Len(t, results, 1, "result size is invalid") - result := results[0].(*packet.BufferData) + require.Len(t, results, 1, "result size is invalid") + result := results[0] require.NoError(t, result.Err) - assert.NotNil(t, result.Buf) + require.NotNil(t, result.Buf) } func TestGeneratorPacketsWithTwoPairs(t *testing.T) { @@ -123,15 +111,15 @@ func TestGeneratorPacketsWithTwoPairs(t *testing.T) { g := NewPacketGenerator(f) out := g.Packets(context.Background(), in) - results := chanToSlice(t, chanBufferDataToGeneric(out), 2) + results := chanToSlice(t, out, 2) - assert.Len(t, results, 2, "result size is invalid") - result1 := results[0].(*packet.BufferData) - result2 := results[1].(*packet.BufferData) + require.Len(t, results, 2, "result size is invalid") + result1 := results[0] + result2 := results[1] require.NoError(t, result1.Err) - assert.NotNil(t, result1.Buf) + require.NotNil(t, result1.Buf) require.NoError(t, result2.Err) - assert.NotNil(t, result2.Buf) + require.NotNil(t, result2.Buf) } func TestMultiGeneratorPacketsWithTwoPairs(t *testing.T) { @@ -155,15 +143,15 @@ func TestMultiGeneratorPacketsWithTwoPairs(t *testing.T) { g := NewPacketMultiGenerator(f, runtime.NumCPU()) out := g.Packets(context.Background(), in) - results := chanToSlice(t, chanBufferDataToGeneric(out), 2) + results := chanToSlice(t, out, 2) - assert.Len(t, results, 2, "result size is invalid") - result1 := results[0].(*packet.BufferData) - result2 := results[1].(*packet.BufferData) + require.Len(t, results, 2, "result size is invalid") + result1 := results[0] + result2 := results[1] require.NoError(t, result1.Err) - assert.NotNil(t, result1.Buf) + require.NotNil(t, result1.Buf) require.NoError(t, result2.Err) - assert.NotNil(t, result2.Buf) + require.NotNil(t, result2.Buf) } func TestGeneratorPacketsReturnsRequestError(t *testing.T) { @@ -178,12 +166,12 @@ func TestGeneratorPacketsReturnsRequestError(t *testing.T) { g := NewPacketGenerator(f) out := g.Packets(context.Background(), in) - results := chanToSlice(t, chanBufferDataToGeneric(out), 1) + results := chanToSlice(t, out, 1) - assert.Len(t, results, 1, "result size is invalid") - result := results[0].(*packet.BufferData) + require.Len(t, results, 1, "result size is invalid") + result := results[0] require.Error(t, result.Err) - assert.Nil(t, result.Buf) + require.Nil(t, result.Buf) } func TestGeneratorPacketsReturnsFillError(t *testing.T) { @@ -203,12 +191,12 @@ func TestGeneratorPacketsReturnsFillError(t *testing.T) { g := NewPacketGenerator(f) out := g.Packets(context.Background(), in) - results := chanToSlice(t, chanBufferDataToGeneric(out), 1) + results := chanToSlice(t, out, 1) - assert.Len(t, results, 1, "result size is invalid") - result := results[0].(*packet.BufferData) + require.Len(t, results, 1, "result size is invalid") + result := results[0] require.Error(t, result.Err) - assert.Nil(t, result.Buf) + require.Nil(t, result.Buf) } func TestMultiGeneratorPacketsReturnsRequestError(t *testing.T) { @@ -223,12 +211,12 @@ func TestMultiGeneratorPacketsReturnsRequestError(t *testing.T) { g := NewPacketMultiGenerator(f, runtime.NumCPU()) out := g.Packets(context.Background(), in) - results := chanToSlice(t, chanBufferDataToGeneric(out), 1) + results := chanToSlice(t, out, 1) - assert.Len(t, results, 1, "result size is invalid") - result := results[0].(*packet.BufferData) + require.Len(t, results, 1, "result size is invalid") + result := results[0] require.Error(t, result.Err) - assert.Nil(t, result.Buf) + require.Nil(t, result.Buf) } func TestMultiGeneratorPacketsReturnsFillError(t *testing.T) { @@ -249,12 +237,12 @@ func TestMultiGeneratorPacketsReturnsFillError(t *testing.T) { g := NewPacketMultiGenerator(f, runtime.NumCPU()) out := g.Packets(context.Background(), in) - results := chanToSlice(t, chanBufferDataToGeneric(out), 1) + results := chanToSlice(t, out, 1) - assert.Len(t, results, 1, "result size is invalid") - result := results[0].(*packet.BufferData) + require.Len(t, results, 1, "result size is invalid") + result := results[0] require.Error(t, result.Err) - assert.Nil(t, result.Buf) + require.Nil(t, result.Buf) } func TestGeneratorPacketsWithTimeout(t *testing.T) { @@ -291,33 +279,33 @@ func TestMultiGeneratorPacketsWithTimeout(t *testing.T) { } } -func TestMergeBufferDataChanEmptyChannels(t *testing.T) { +func TestMergeChannelsBufferDataEmptyChannels(t *testing.T) { t.Parallel() c1 := make(chan *packet.BufferData) close(c1) c2 := make(chan *packet.BufferData) close(c2) - out := MergeBufferDataChan(context.Background(), c1, c2) + out := mergeChannels(context.Background(), 200, c1, c2) - result := chanToSlice(t, chanBufferDataToGeneric(out), 0) - assert.Empty(t, result, "result slice is not empty") + result := chanToSlice(t, out, 0) + require.Empty(t, result, "result slice is not empty") } -func TestMergeBufferDataChanOneElementAndEmptyChannel(t *testing.T) { +func TestMergeChannelsBufferDataOneElementAndEmptyChannel(t *testing.T) { t.Parallel() c1 := make(chan *packet.BufferData, 1) c1 <- &packet.BufferData{} close(c1) c2 := make(chan *packet.BufferData) close(c2) - out := MergeBufferDataChan(context.Background(), c1, c2) + out := mergeChannels(context.Background(), 200, c1, c2) - result := chanToSlice(t, chanBufferDataToGeneric(out), 1) - assert.Len(t, result, 1, "result slice size is invalid") - assert.NotNil(t, result[0]) + result := chanToSlice(t, out, 1) + require.Len(t, result, 1, "result slice size is invalid") + require.NotNil(t, result[0]) } -func TestMergeBufferDataChanTwoElements(t *testing.T) { +func TestMergeChannelsBufferDataTwoElements(t *testing.T) { t.Parallel() c1 := make(chan *packet.BufferData, 1) c1 <- &packet.BufferData{} @@ -325,15 +313,15 @@ func TestMergeBufferDataChanTwoElements(t *testing.T) { c2 := make(chan *packet.BufferData, 1) c2 <- &packet.BufferData{} close(c2) - out := MergeBufferDataChan(context.Background(), c1, c2) + out := mergeChannels(context.Background(), 200, c1, c2) - result := chanToSlice(t, chanBufferDataToGeneric(out), 2) - assert.Len(t, result, 2, "result slice size is invalid") - assert.NotNil(t, result[0]) - assert.NotNil(t, result[1]) + result := chanToSlice(t, out, 2) + require.Len(t, result, 2, "result slice size is invalid") + require.NotNil(t, result[0]) + require.NotNil(t, result[1]) } -func TestMergeBufferDataChanContextExit(t *testing.T) { +func TestMergeChannelsBufferDataContextExit(t *testing.T) { t.Parallel() c1 := make(chan *packet.BufferData) defer close(c1) @@ -342,8 +330,8 @@ func TestMergeBufferDataChanContextExit(t *testing.T) { ctx, cancel := context.WithTimeout(context.Background(), 1*time.Millisecond) defer cancel() - out := MergeBufferDataChan(ctx, c1, c2) + out := mergeChannels(ctx, 200, c1, c2) - result := chanToSlice(t, chanBufferDataToGeneric(out), 0) - assert.Empty(t, result, "result slice is not empty") + result := chanToSlice(t, out, 0) + require.Empty(t, result, "result slice is not empty") } diff --git a/pkg/scan/mock_request_test.go b/pkg/scan/mock_request_test.go index dbc013f..182b9c5 100644 --- a/pkg/scan/mock_request_test.go +++ b/pkg/scan/mock_request_test.go @@ -42,10 +42,10 @@ func (m *MockPortGenerator) EXPECT() *MockPortGeneratorMockRecorder { } // Ports mocks base method. -func (m *MockPortGenerator) Ports(ctx context.Context, r *Range) (<-chan PortGetter, error) { +func (m *MockPortGenerator) Ports(ctx context.Context, r *Range) (<-chan GeneratorResult[uint16], error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "Ports", ctx, r) - ret0, _ := ret[0].(<-chan PortGetter) + ret0, _ := ret[0].(<-chan GeneratorResult[uint16]) ret1, _ := ret[1].(error) return ret0, ret1 } @@ -81,10 +81,10 @@ func (m *MockIPGenerator) EXPECT() *MockIPGeneratorMockRecorder { } // IPs mocks base method. -func (m *MockIPGenerator) IPs(ctx context.Context, r *Range) (<-chan IPGetter, error) { +func (m *MockIPGenerator) IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[net.IP], error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "IPs", ctx, r) - ret0, _ := ret[0].(<-chan IPGetter) + ret0, _ := ret[0].(<-chan GeneratorResult[net.IP]) ret1, _ := ret[1].(error) return ret0, ret1 } diff --git a/pkg/scan/mock_sendreceiver_test.go b/pkg/scan/mock_sendreceiver_test.go index 0434634..7272a9b 100644 --- a/pkg/scan/mock_sendreceiver_test.go +++ b/pkg/scan/mock_sendreceiver_test.go @@ -42,10 +42,10 @@ func (m *MockSender) EXPECT() *MockSenderMockRecorder { } // SendPackets mocks base method. -func (m *MockSender) SendPackets(ctx context.Context, in <-chan *packet.BufferData) (<-chan any, <-chan error) { +func (m *MockSender) SendPackets(ctx context.Context, in <-chan *packet.BufferData) (<-chan struct{}, <-chan error) { m.ctrl.T.Helper() ret := m.ctrl.Call(m, "SendPackets", ctx, in) - ret0, _ := ret[0].(<-chan any) + ret0, _ := ret[0].(<-chan struct{}) ret1, _ := ret[1].(<-chan error) return ret0, ret1 } diff --git a/pkg/scan/request.go b/pkg/scan/request.go index 0b695ac..72a464c 100644 --- a/pkg/scan/request.go +++ b/pkg/scan/request.go @@ -31,26 +31,18 @@ type Request struct { Err error } -type PortGetter interface { - GetPort() (uint16, error) -} - -type WrapPort uint16 - -func (p WrapPort) GetPort() (uint16, error) { - return uint16(p), nil -} - -type portError struct { - error -} - -func (err *portError) GetPort() (uint16, error) { - return 0, err +// GeneratorResult contains one asynchronously generated value or its error. +type GeneratorResult[T any] struct { + // Value contains the generated value when Err is nil. + Value T + // Err contains an error for this generated item. + Err error } +// PortGenerator produces ports for a scan range. type PortGenerator interface { - Ports(ctx context.Context, r *Range) (<-chan PortGetter, error) + // Ports generates the ports described by r until completion or context cancellation. + Ports(ctx context.Context, r *Range) (<-chan GeneratorResult[uint16], error) } func NewPortGenerator() PortGenerator { @@ -59,22 +51,28 @@ func NewPortGenerator() PortGenerator { type portGenerator struct{} -func (*portGenerator) Ports(ctx context.Context, r *Range) (<-chan PortGetter, error) { +func (*portGenerator) Ports(ctx context.Context, r *Range) (<-chan GeneratorResult[uint16], error) { if err := validatePorts(r.Ports); err != nil { return nil, err } - out := make(chan PortGetter, 100) + out := make(chan GeneratorResult[uint16], 100) go func() { defer close(out) for _, portRange := range r.Ports { it, err := newRangeIterator(int64(portRange.EndPort) - int64(portRange.StartPort) + 1) if err != nil { - writePort(ctx, out, &portError{err}) + if !sendContext(ctx, out, GeneratorResult[uint16]{Err: err}) { + return + } continue } basePort := int64(portRange.StartPort) - 1 for { - writePort(ctx, out, WrapPort(basePort+it.Int().Int64())) + if !sendContext(ctx, out, GeneratorResult[uint16]{ + Value: uint16(basePort + it.Int().Int64()), + }) { + return + } if !it.Next() { break } @@ -84,14 +82,6 @@ func (*portGenerator) Ports(ctx context.Context, r *Range) (<-chan PortGetter, e return out, nil } -func writePort(ctx context.Context, out chan<- PortGetter, port PortGetter) { - select { - case <-ctx.Done(): - return - case out <- port: - } -} - func validatePorts(ports []*PortRange) error { if len(ports) == 0 { return ErrPortRange @@ -104,18 +94,10 @@ func validatePorts(ports []*PortRange) error { return nil } -type IPGetter interface { - GetIP() (net.IP, error) -} - -type WrapIP net.IP - -func (i WrapIP) GetIP() (net.IP, error) { - return net.IP(i), nil -} - +// IPGenerator produces IP addresses for a scan range. type IPGenerator interface { - IPs(ctx context.Context, r *Range) (<-chan IPGetter, error) + // IPs generates the IP addresses described by r until completion or context cancellation. + IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[net.IP], error) } func NewIPGenerator() IPGenerator { @@ -124,7 +106,7 @@ func NewIPGenerator() IPGenerator { type ipGenerator struct{} -func (*ipGenerator) IPs(ctx context.Context, r *Range) (<-chan IPGetter, error) { +func (*ipGenerator) IPs(ctx context.Context, r *Range) (<-chan GeneratorResult[net.IP], error) { if r.DstSubnet == nil { return nil, ErrSubnet } @@ -138,7 +120,7 @@ func (*ipGenerator) IPs(ctx context.Context, r *Range) (<-chan IPGetter, error) baseIP := big.NewInt(0).SetBytes(ipnet.IP.Mask(ipnet.Mask)) baseIP.Sub(baseIP, big.NewInt(1)) - out := make(chan IPGetter, 100) + out := make(chan GeneratorResult[net.IP], 100) go func() { defer close(out) for { @@ -148,7 +130,9 @@ func (*ipGenerator) IPs(ctx context.Context, r *Range) (<-chan IPGetter, error) ipaddr := baseIP.FillBytes(make([]byte, 4)) baseIP.Sub(baseIP, i) - writeIP(ctx, out, WrapIP(ipaddr)) + if !sendContext(ctx, out, GeneratorResult[net.IP]{Value: ipaddr}) { + return + } if !it.Next() { return @@ -184,34 +168,27 @@ func (rg *ipPortGenerator) GenerateRequests(ctx context.Context, r *Range) (<-ch go func() { defer close(out) for p := range ports { - port, err := p.GetPort() - if err != nil { - writeRequest(ctx, out, &Request{Err: err}) + if p.Err != nil { + if !sendContext(ctx, out, &Request{Err: p.Err}) { + return + } continue } for ipaddr := range ips { - dstip, err := ipaddr.GetIP() - writeRequest(ctx, out, &Request{ + if !sendContext(ctx, out, &Request{ SrcIP: r.SrcIP, SrcMAC: r.SrcMAC, - DstIP: dstip, DstPort: port, Err: err}) + DstIP: ipaddr.Value, DstPort: p.Value, Err: ipaddr.Err}) { + return + } } if ips, err = rg.ipgen.IPs(ctx, r); err != nil { - writeRequest(ctx, out, &Request{Err: err}) + sendContext(ctx, out, &Request{Err: err}) return } } }() return out, nil } - -func writeRequest(ctx context.Context, out chan<- *Request, request *Request) { - select { - case <-ctx.Done(): - return - case out <- request: - } -} - func NewIPRequestGenerator(ipgen IPGenerator) RequestGenerator { return &ipRequestGenerator{ipgen} } @@ -229,11 +206,11 @@ func (rg *ipRequestGenerator) GenerateRequests(ctx context.Context, r *Range) (< go func() { defer close(out) for ipaddr := range ips { - dstip, err := ipaddr.GetIP() - writeRequest(ctx, out, &Request{ - SrcIP: r.SrcIP, SrcMAC: r.SrcMAC, DstIP: dstip, - Err: err, - }) + if !sendContext(ctx, out, &Request{ + SrcIP: r.SrcIP, SrcMAC: r.SrcMAC, DstIP: ipaddr.Value, + Err: ipaddr.Err}) { + return + } } }() return out, nil @@ -270,23 +247,29 @@ func (rg *fileIPPortGenerator) GenerateRequests(ctx context.Context, r *Range) ( entry.IP = "" entry.Port = 0 if err := entry.UnmarshalJSON(scanner.Bytes()); err != nil { - writeRequest(ctx, out, &Request{Err: ErrJSON}) + sendContext(ctx, out, &Request{Err: ErrJSON}) return } ip := net.ParseIP(entry.IP) if ip == nil { - writeRequest(ctx, out, &Request{Err: ErrIP}) + if !sendContext(ctx, out, &Request{Err: ErrIP}) { + return + } continue } if !isValidPort(entry.Port) { - writeRequest(ctx, out, &Request{Err: ErrPort}) + if !sendContext(ctx, out, &Request{Err: ErrPort}) { + return + } continue } - writeRequest(ctx, out, &Request{ - SrcIP: r.SrcIP, SrcMAC: r.SrcMAC, DstIP: ip, DstPort: uint16(entry.Port)}) + if !sendContext(ctx, out, &Request{ + SrcIP: r.SrcIP, SrcMAC: r.SrcMAC, DstIP: ip, DstPort: uint16(entry.Port)}) { + return + } } if err = scanner.Err(); err != nil { - writeRequest(ctx, out, &Request{Err: err}) + sendContext(ctx, out, &Request{Err: err}) } }() return out, nil @@ -296,14 +279,6 @@ func isValidPort(port int) bool { return port > 0 && port <= 0xFFFF } -type ipError struct { - error -} - -func (err *ipError) GetIP() (net.IP, error) { - return nil, err -} - type fileIPGenerator struct { openFile OpenFileFunc } @@ -312,44 +287,39 @@ func NewFileIPGenerator(openFile OpenFileFunc) IPGenerator { return &fileIPGenerator{openFile} } -func (g *fileIPGenerator) IPs(ctx context.Context, _ *Range) (<-chan IPGetter, error) { +func (g *fileIPGenerator) IPs(ctx context.Context, _ *Range) (<-chan GeneratorResult[net.IP], error) { input, err := g.openFile() if err != nil { return nil, err } - out := make(chan IPGetter) + out := make(chan GeneratorResult[net.IP]) go func() { defer close(out) defer input.Close() scanner := bufio.NewScanner(input) var entry IPPort for scanner.Scan() { + entry = IPPort{} if err := entry.UnmarshalJSON(scanner.Bytes()); err != nil { - writeIP(ctx, out, &ipError{error: ErrJSON}) + sendContext(ctx, out, GeneratorResult[net.IP]{Err: ErrJSON}) return } ip := net.ParseIP(entry.IP) if ip == nil { - writeIP(ctx, out, &ipError{error: ErrIP}) + sendContext(ctx, out, GeneratorResult[net.IP]{Err: ErrIP}) + return + } + if !sendContext(ctx, out, GeneratorResult[net.IP]{Value: ip}) { return } - writeIP(ctx, out, WrapIP(ip)) } if err = scanner.Err(); err != nil { - writeIP(ctx, out, &ipError{error: err}) + sendContext(ctx, out, GeneratorResult[net.IP]{Err: err}) } }() return out, nil } -func writeIP(ctx context.Context, out chan<- IPGetter, ip IPGetter) { - select { - case <-ctx.Done(): - return - case out <- ip: - } -} - type liveRequestGenerator struct { delegate RequestGenerator rescanTimeout time.Duration @@ -371,14 +341,20 @@ func (rg *liveRequestGenerator) GenerateRequests(ctx context.Context, r *Range) var ok bool for { if request, ok = readRequest(ctx, requests); ok { - writeRequest(ctx, out, request) + if !sendContext(ctx, out, request) { + return + } continue } select { case <-ctx.Done(): return case <-time.After(rg.rescanTimeout): - requests, _ = rg.delegate.GenerateRequests(ctx, r) + requests, err = rg.delegate.GenerateRequests(ctx, r) + if err != nil { + sendContext(ctx, out, &Request{Err: err}) + return + } } } }() @@ -423,13 +399,17 @@ func (rg *filterIPRequestGenerator) GenerateRequests(ctx context.Context, r *Ran contains, err := rg.excludeIPs.Contains(request.DstIP) if err != nil { request.Err = err - writeRequest(ctx, out, request) + if !sendContext(ctx, out, request) { + return + } continue } if contains { continue } - writeRequest(ctx, out, request) + if !sendContext(ctx, out, request) { + return + } } }() return out, nil diff --git a/pkg/scan/request_test.go b/pkg/scan/request_test.go index d0312fe..6fffe93 100644 --- a/pkg/scan/request_test.go +++ b/pkg/scan/request_test.go @@ -12,7 +12,6 @@ import ( "testing" "time" - "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "go.uber.org/mock/gomock" ) @@ -83,15 +82,26 @@ func withError(err error) scanRequestOption { } } -func chanPortToGeneric(in <-chan PortGetter) <-chan interface{} { - out := make(chan interface{}, cap(in)) - go func() { - defer close(out) - for i := range in { - out <- i - } - }() - return out +func TestGeneratorResult(t *testing.T) { + t.Parallel() + + err := errors.New("generated value error") + result := GeneratorResult[int]{Value: 42, Err: err} + + require.Equal(t, 42, result.Value) + require.ErrorIs(t, result.Err, err) +} + +func generatedPort(value uint16) GeneratorResult[uint16] { + return GeneratorResult[uint16]{Value: value} +} + +func generatedIP(value net.IP) GeneratorResult[net.IP] { + return GeneratorResult[net.IP]{Value: value} +} + +func generationError[T any](err error) GeneratorResult[T] { + return GeneratorResult[T]{Err: err} } func TestPortGenerator(t *testing.T) { @@ -100,7 +110,7 @@ func TestPortGenerator(t *testing.T) { tests := []struct { name string scanRange *Range - expected []interface{} + expected []GeneratorResult[uint16] err bool }{ { @@ -140,7 +150,7 @@ func TestPortGenerator(t *testing.T) { EndPort: 22, }, })), - expected: []interface{}{WrapPort(22)}, + expected: []GeneratorResult[uint16]{generatedPort(22)}, }, { name: "TwoPorts", @@ -150,7 +160,7 @@ func TestPortGenerator(t *testing.T) { EndPort: 23, }, })), - expected: []interface{}{WrapPort(22), WrapPort(23)}, + expected: []GeneratorResult[uint16]{generatedPort(22), generatedPort(23)}, }, { name: "ThreePorts", @@ -160,7 +170,7 @@ func TestPortGenerator(t *testing.T) { EndPort: 27, }, })), - expected: []interface{}{WrapPort(25), WrapPort(26), WrapPort(27)}, + expected: []GeneratorResult[uint16]{generatedPort(25), generatedPort(26), generatedPort(27)}, }, { name: "OnePortOverflow", @@ -170,7 +180,7 @@ func TestPortGenerator(t *testing.T) { EndPort: 65535, }, })), - expected: []interface{}{WrapPort(65535)}, + expected: []GeneratorResult[uint16]{generatedPort(65535)}, }, { name: "TwoRangesOnePort", @@ -184,7 +194,7 @@ func TestPortGenerator(t *testing.T) { EndPort: 27, }, })), - expected: []interface{}{WrapPort(25), WrapPort(27)}, + expected: []GeneratorResult[uint16]{generatedPort(25), generatedPort(27)}, }, { name: "TwoRangesTwoPorts", @@ -198,8 +208,8 @@ func TestPortGenerator(t *testing.T) { EndPort: 27, }, })), - expected: []interface{}{WrapPort(20), WrapPort(21), WrapPort(23), - WrapPort(24), WrapPort(25), WrapPort(26), WrapPort(27)}, + expected: []GeneratorResult[uint16]{generatedPort(20), generatedPort(21), generatedPort(23), + generatedPort(24), generatedPort(25), generatedPort(26), generatedPort(27)}, }, { name: "ZeroPort", @@ -209,7 +219,7 @@ func TestPortGenerator(t *testing.T) { EndPort: 1, }, })), - expected: []interface{}{WrapPort(0), WrapPort(1)}, + expected: []GeneratorResult[uint16]{generatedPort(0), generatedPort(1)}, }, } @@ -218,77 +228,49 @@ func TestPortGenerator(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - portgen := NewPortGenerator() - ports, err := portgen.Ports(context.Background(), tt.scanRange) - if tt.err { - assert.Error(t, err) - return - } - if !assert.NoError(t, err) { - return - } - result := collectInterfaces(chanPortToGeneric(ports)) - sort.Slice(result, func(i, j int) bool { - return uint16(result[i].(WrapPort)) < uint16(result[j].(WrapPort)) - }) - assert.Equal(t, tt.expected, result) - }() - waitDone(t, done) + portgen := NewPortGenerator() + ports, err := portgen.Ports(context.Background(), tt.scanRange) + if tt.err { + require.Error(t, err) + return + } + require.NoError(t, err) + result := collectChannel(ports) + sort.Slice(result, func(i, j int) bool { + return result[i].Value < result[j].Value + }) + require.Equal(t, tt.expected, result) }) } } func TestPortGeneratorFullRange(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - portgen := NewPortGenerator() - ports, err := portgen.Ports(context.Background(), newScanRange(withPorts([]*PortRange{ - { - StartPort: 1, - EndPort: 65535, - }, - }))) - if !assert.NoError(t, err) { - return - } - - bitset := big.NewInt(0) - cnt := 0 - for p := range ports { - cnt++ - port, err := p.GetPort() - if !assert.NoError(t, err) { - return - } - i := int(port) - if bitset.Bit(i) == 1 { - assert.Fail(t, "number has already been visited", "number %d", i) - } - bitset.SetBit(bitset, i, 1) - } - for i := 1; i <= 65535; i++ { - assert.Equal(t, uint(1), bitset.Bit(i), - "number %d is not visited", i) - } - assert.Equal(t, 65535, cnt, "count is not valid") - }() - waitDone(t, done) -} + portgen := NewPortGenerator() + ports, err := portgen.Ports(context.Background(), newScanRange(withPorts([]*PortRange{ + { + StartPort: 1, + EndPort: 65535, + }, + }))) + require.NoError(t, err) -func chanIPToGeneric(in <-chan IPGetter) <-chan interface{} { - out := make(chan interface{}, cap(in)) - go func() { - defer close(out) - for i := range in { - out <- i + bitset := big.NewInt(0) + cnt := 0 + for p := range ports { + cnt++ + require.NoError(t, p.Err) + i := int(p.Value) + if bitset.Bit(i) == 1 { + require.Fail(t, "number has already been visited", "number %d", i) } - }() - return out + bitset.SetBit(bitset, i, 1) + } + for i := 1; i <= 65535; i++ { + require.Equal(t, uint(1), bitset.Bit(i), + "number %d is not visited", i) + } + require.Equal(t, 65535, cnt, "count is not valid") } func TestIPGenerator(t *testing.T) { @@ -297,7 +279,7 @@ func TestIPGenerator(t *testing.T) { tests := []struct { name string scanRange *Range - expected []interface{} + expected []GeneratorResult[net.IP] err bool }{ { @@ -310,8 +292,8 @@ func TestIPGenerator(t *testing.T) { scanRange: newScanRange( withSubnet(&net.IPNet{IP: net.IPv4(192, 168, 0, 1), Mask: net.CIDRMask(32, 32)}), ), - expected: []interface{}{ - WrapIP(net.IPv4(192, 168, 0, 1).To4()), + expected: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(192, 168, 0, 1).To4()), }, }, { @@ -319,9 +301,9 @@ func TestIPGenerator(t *testing.T) { scanRange: newScanRange( withSubnet(&net.IPNet{IP: net.IPv4(1, 0, 0, 1), Mask: net.CIDRMask(31, 32)}), ), - expected: []interface{}{ - WrapIP(net.IPv4(1, 0, 0, 0).To4()), - WrapIP(net.IPv4(1, 0, 0, 1).To4()), + expected: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(1, 0, 0, 0).To4()), + generatedIP(net.IPv4(1, 0, 0, 1).To4()), }, }, { @@ -329,11 +311,11 @@ func TestIPGenerator(t *testing.T) { scanRange: newScanRange( withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 1), Mask: net.CIDRMask(30, 32)}), ), - expected: []interface{}{ - WrapIP(net.IPv4(10, 0, 0, 0).To4()), - WrapIP(net.IPv4(10, 0, 0, 1).To4()), - WrapIP(net.IPv4(10, 0, 0, 2).To4()), - WrapIP(net.IPv4(10, 0, 0, 3).To4()), + expected: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(10, 0, 0, 0).To4()), + generatedIP(net.IPv4(10, 0, 0, 1).To4()), + generatedIP(net.IPv4(10, 0, 0, 2).To4()), + generatedIP(net.IPv4(10, 0, 0, 3).To4()), }, }, } @@ -343,42 +325,24 @@ func TestIPGenerator(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - ipgen := NewIPGenerator() - ips, err := ipgen.IPs(context.Background(), tt.scanRange) - if tt.err { - assert.Error(t, err) - return - } - if !assert.NoError(t, err) { - return - } - result := collectInterfaces(chanIPToGeneric(ips)) - sort.Slice(result, func(i, j int) bool { - return bytes.Compare(result[i].(WrapIP), result[j].(WrapIP)) < 1 - }) - assert.Equal(t, tt.expected, result) - }() - waitDone(t, done) + ipgen := NewIPGenerator() + ips, err := ipgen.IPs(context.Background(), tt.scanRange) + if tt.err { + require.Error(t, err) + return + } + require.NoError(t, err) + result := collectChannel(ips) + sort.Slice(result, func(i, j int) bool { + return bytes.Compare(result[i].Value, result[j].Value) < 1 + }) + require.Equal(t, tt.expected, result) }) } } -func chanPairToGeneric(in <-chan *Request) <-chan interface{} { - out := make(chan interface{}, cap(in)) - go func() { - defer close(out) - for i := range in { - out <- i - } - }() - return out -} - -func collectInterfaces(in <-chan interface{}) []interface{} { - var out []interface{} +func collectChannel[T any](in <-chan T) []T { + var out []T for v := range in { out = append(out, v) } @@ -390,36 +354,36 @@ func TestIPPortGenerator(t *testing.T) { tests := []struct { name string - ips []IPGetter - ports []PortGetter - expected []interface{} + ips []GeneratorResult[net.IP] + ports []GeneratorResult[uint16] + expected []*Request }{ { name: "OneIpOnePort", - ips: []IPGetter{WrapIP(net.IPv4(192, 168, 0, 1))}, - ports: []PortGetter{WrapPort(888)}, - expected: []interface{}{ + ips: []GeneratorResult[net.IP]{generatedIP(net.IPv4(192, 168, 0, 1))}, + ports: []GeneratorResult[uint16]{generatedPort(888)}, + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), }, }, { name: "OneIpTwoPorts", - ips: []IPGetter{WrapIP(net.IPv4(192, 168, 0, 1))}, - ports: []PortGetter{WrapPort(888), WrapPort(889)}, - expected: []interface{}{ + ips: []GeneratorResult[net.IP]{generatedIP(net.IPv4(192, 168, 0, 1))}, + ports: []GeneratorResult[uint16]{generatedPort(888), generatedPort(889)}, + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(889)), }, }, { name: "ThreeIpsOnePort", - ips: []IPGetter{ - WrapIP(net.IPv4(192, 168, 0, 1)), - WrapIP(net.IPv4(192, 168, 0, 2)), - WrapIP(net.IPv4(192, 168, 0, 3)), + ips: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(192, 168, 0, 1)), + generatedIP(net.IPv4(192, 168, 0, 2)), + generatedIP(net.IPv4(192, 168, 0, 3)), }, - ports: []PortGetter{WrapPort(888)}, - expected: []interface{}{ + ports: []GeneratorResult[uint16]{generatedPort(888)}, + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), newScanRequest(withDstIP(net.IPv4(192, 168, 0, 2)), withDstPort(888)), newScanRequest(withDstIP(net.IPv4(192, 168, 0, 3)), withDstPort(888)), @@ -427,12 +391,12 @@ func TestIPPortGenerator(t *testing.T) { }, { name: "TwoIpsTwoPorts", - ips: []IPGetter{ - WrapIP(net.IPv4(192, 168, 0, 1)), - WrapIP(net.IPv4(192, 168, 0, 2)), + ips: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(192, 168, 0, 1)), + generatedIP(net.IPv4(192, 168, 0, 2)), }, - ports: []PortGetter{WrapPort(888), WrapPort(889)}, - expected: []interface{}{ + ports: []GeneratorResult[uint16]{generatedPort(888), generatedPort(889)}, + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), newScanRequest(withDstIP(net.IPv4(192, 168, 0, 2)), withDstPort(888)), newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(889)), @@ -441,33 +405,33 @@ func TestIPPortGenerator(t *testing.T) { }, { name: "IPError", - ips: []IPGetter{ - &ipError{errors.New("ip error")}, + ips: []GeneratorResult[net.IP]{ + generationError[net.IP](errors.New("ip error")), }, - ports: []PortGetter{WrapPort(888)}, - expected: []interface{}{ - newScanRequest(withDstIP(nil), withDstPort(888), withError(&ipError{errors.New("ip error")})), + ports: []GeneratorResult[uint16]{generatedPort(888)}, + expected: []*Request{ + newScanRequest(withDstIP(nil), withDstPort(888), withError(errors.New("ip error"))), }, }, { name: "PortError", - ips: []IPGetter{WrapIP(net.IPv4(192, 168, 0, 1))}, - ports: []PortGetter{ - &portError{errors.New("port error")}, + ips: []GeneratorResult[net.IP]{generatedIP(net.IPv4(192, 168, 0, 1))}, + ports: []GeneratorResult[uint16]{ + generationError[uint16](errors.New("port error")), }, - expected: []interface{}{ - &Request{Err: &portError{errors.New("port error")}}, + expected: []*Request{ + {Err: errors.New("port error")}, }, }, { name: "ValidPortAfterPortError", - ips: []IPGetter{WrapIP(net.IPv4(192, 168, 0, 1))}, - ports: []PortGetter{ - &portError{errors.New("port error")}, - WrapPort(888), + ips: []GeneratorResult[net.IP]{generatedIP(net.IPv4(192, 168, 0, 1))}, + ports: []GeneratorResult[uint16]{ + generationError[uint16](errors.New("port error")), + generatedPort(888), }, - expected: []interface{}{ - &Request{Err: &portError{errors.New("port error")}}, + expected: []*Request{ + {Err: errors.New("port error")}, newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1)), withDstPort(888)), }, }, @@ -478,43 +442,35 @@ func TestIPPortGenerator(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - ctrl := gomock.NewController(t) - ipgen := NewMockIPGenerator(ctrl) - - ctx := context.Background() - scanRange := newScanRange() - ipgen.EXPECT().IPs(ctx, scanRange). - DoAndReturn(func(ctx context.Context, r *Range) (<-chan IPGetter, error) { - ips := make(chan IPGetter, len(tt.ips)) - for _, ip := range tt.ips { - ips <- ip - } - close(ips) - return ips, nil - }).AnyTimes() - - ports := make(chan PortGetter, len(tt.ports)) - for _, port := range tt.ports { - ports <- port - } - close(ports) + ctrl := gomock.NewController(t) + ipgen := NewMockIPGenerator(ctrl) - portgen := NewMockPortGenerator(ctrl) - portgen.EXPECT().Ports(ctx, scanRange).Return(ports, nil) + ctx := context.Background() + scanRange := newScanRange() + ipgen.EXPECT().IPs(ctx, scanRange). + DoAndReturn(func(ctx context.Context, r *Range) (<-chan GeneratorResult[net.IP], error) { + ips := make(chan GeneratorResult[net.IP], len(tt.ips)) + for _, ip := range tt.ips { + ips <- ip + } + close(ips) + return ips, nil + }).AnyTimes() - reqgen := NewIPPortGenerator(ipgen, portgen) - pairs, err := reqgen.GenerateRequests(ctx, scanRange) - if !assert.NoError(t, err) { - return - } - result := collectInterfaces(chanPairToGeneric(pairs)) - assert.Equal(t, tt.expected, result) - }() - waitDone(t, done) + ports := make(chan GeneratorResult[uint16], len(tt.ports)) + for _, port := range tt.ports { + ports <- port + } + close(ports) + + portgen := NewMockPortGenerator(ctrl) + portgen.EXPECT().Ports(ctx, scanRange).Return(ports, nil) + + reqgen := NewIPPortGenerator(ipgen, portgen) + pairs, err := reqgen.GenerateRequests(ctx, scanRange) + require.NoError(t, err) + result := collectChannel(pairs) + require.Equal(t, tt.expected, result) }) } } @@ -542,25 +498,19 @@ func TestIPPortGeneratorError(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) + ctrl := gomock.NewController(t) + ipgen := NewMockIPGenerator(ctrl) - ctrl := gomock.NewController(t) - ipgen := NewMockIPGenerator(ctrl) + ctx := context.Background() + scanRange := newScanRange() + ipgen.EXPECT().IPs(ctx, scanRange).Return(nil, tt.ipsError).AnyTimes() - ctx := context.Background() - scanRange := newScanRange() - ipgen.EXPECT().IPs(ctx, scanRange).Return(nil, tt.ipsError).AnyTimes() + portgen := NewMockPortGenerator(ctrl) + portgen.EXPECT().Ports(ctx, scanRange).Return(nil, tt.portsError).AnyTimes() - portgen := NewMockPortGenerator(ctrl) - portgen.EXPECT().Ports(ctx, scanRange).Return(nil, tt.portsError).AnyTimes() - - reqgen := NewIPPortGenerator(ipgen, portgen) - _, err := reqgen.GenerateRequests(ctx, scanRange) - assert.Error(t, err) - }() - waitDone(t, done) + reqgen := NewIPPortGenerator(ipgen, portgen) + _, err := reqgen.GenerateRequests(ctx, scanRange) + require.Error(t, err) }) } } @@ -571,7 +521,7 @@ func TestIPRequestGenerator(t *testing.T) { tests := []struct { name string input *Range - expected []interface{} + expected []*Request err bool }{ { @@ -584,7 +534,7 @@ func TestIPRequestGenerator(t *testing.T) { input: newScanRange( withSubnet(&net.IPNet{IP: net.IPv4(192, 168, 0, 1).To4(), Mask: net.CIDRMask(32, 32)}), ), - expected: []interface{}{ + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1).To4())), }, }, @@ -593,7 +543,7 @@ func TestIPRequestGenerator(t *testing.T) { input: newScanRange( withSubnet(&net.IPNet{IP: net.IPv4(192, 168, 0, 1).To4(), Mask: net.CIDRMask(31, 32)}), ), - expected: []interface{}{ + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(192, 168, 0, 0).To4())), newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1).To4())), }, @@ -603,7 +553,7 @@ func TestIPRequestGenerator(t *testing.T) { input: newScanRange( withSubnet(&net.IPNet{IP: net.IPv4(192, 168, 0, 1).To4(), Mask: net.CIDRMask(30, 32)}), ), - expected: []interface{}{ + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(192, 168, 0, 0).To4())), newScanRequest(withDstIP(net.IPv4(192, 168, 0, 1).To4())), newScanRequest(withDstIP(net.IPv4(192, 168, 0, 2).To4())), @@ -617,28 +567,20 @@ func TestIPRequestGenerator(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - reqgen := NewIPRequestGenerator(NewIPGenerator()) - pairs, err := reqgen.GenerateRequests(context.Background(), tt.input) - if tt.err { - assert.Error(t, err) - return - } - if !assert.NoError(t, err) { - return - } - result := collectInterfaces(chanPairToGeneric(pairs)) - sort.Slice(result, func(i, j int) bool { - return bytes.Compare( - result[i].(*Request).DstIP, - result[j].(*Request).DstIP) < 1 - }) - assert.Equal(t, tt.expected, result) - }() - waitDone(t, done) + reqgen := NewIPRequestGenerator(NewIPGenerator()) + pairs, err := reqgen.GenerateRequests(context.Background(), tt.input) + if tt.err { + require.Error(t, err) + return + } + require.NoError(t, err) + result := collectChannel(pairs) + sort.Slice(result, func(i, j int) bool { + return bytes.Compare( + result[i].DstIP, + result[j].DstIP) < 1 + }) + require.Equal(t, tt.expected, result) }) } } @@ -660,20 +602,20 @@ func TestFileIPPortGenerator(t *testing.T) { name string input string scanRange *Range - expected []interface{} + expected []*Request }{ { name: "OneIPPort", input: `{"ip":"192.168.0.1","port":888}`, - expected: []interface{}{ - &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + expected: []*Request{ + {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, }, }, { name: "OneIPPortWithUnknownField", input: `{"ip":"192.168.0.1","port":888,"abc":"field"}`, - expected: []interface{}{ - &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + expected: []*Request{ + {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, }, }, { @@ -682,16 +624,16 @@ func TestFileIPPortGenerator(t *testing.T) { `{"ip":"192.168.0.1","port":888}`, `{"ip":"192.168.0.2","port":222}`, }, "\n"), - expected: []interface{}{ - &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, - &Request{DstIP: net.IPv4(192, 168, 0, 2), DstPort: 222}, + expected: []*Request{ + {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {DstIP: net.IPv4(192, 168, 0, 2), DstPort: 222}, }, }, { name: "InvalidJSON", input: `{"ip":"192`, - expected: []interface{}{ - &Request{Err: ErrJSON}, + expected: []*Request{ + {Err: ErrJSON}, }, }, { @@ -700,9 +642,9 @@ func TestFileIPPortGenerator(t *testing.T) { `{"ip":"192.168.0.1","port":888}`, `{"ip":"192`, }, "\n"), - expected: []interface{}{ - &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, - &Request{Err: ErrJSON}, + expected: []*Request{ + {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {Err: ErrJSON}, }, }, { @@ -712,23 +654,23 @@ func TestFileIPPortGenerator(t *testing.T) { `{"ip":"192`, `{"ip":"192.168.0.3","port":888}`, }, "\n"), - expected: []interface{}{ - &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, - &Request{Err: ErrJSON}, + expected: []*Request{ + {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {Err: ErrJSON}, }, }, { name: "InvalidIP", input: `{"ip":"192.168.0.1111","port":888}`, - expected: []interface{}{ - &Request{Err: ErrIP}, + expected: []*Request{ + {Err: ErrIP}, }, }, { name: "InvalidPort", input: `{"ip":"192.168.0.1","port":88888}`, - expected: []interface{}{ - &Request{Err: ErrPort}, + expected: []*Request{ + {Err: ErrPort}, }, }, { @@ -737,9 +679,9 @@ func TestFileIPPortGenerator(t *testing.T) { `{"ip":"192.168.0.1","port":888}`, `{"ip":"192.168.0.3"}`, }, "\n"), - expected: []interface{}{ - &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, - &Request{Err: ErrPort}, + expected: []*Request{ + {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {Err: ErrPort}, }, }, { @@ -748,9 +690,9 @@ func TestFileIPPortGenerator(t *testing.T) { `{"ip":"192.168.0.1","port":888}`, `{"port":888}`, }, "\n"), - expected: []interface{}{ - &Request{DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, - &Request{Err: ErrIP}, + expected: []*Request{ + {DstIP: net.IPv4(192, 168, 0, 1), DstPort: 888}, + {Err: ErrIP}, }, }, { @@ -760,8 +702,8 @@ func TestFileIPPortGenerator(t *testing.T) { SrcIP: net.IPv4(192, 168, 0, 3), SrcMAC: net.HardwareAddr{0x01, 0x02, 0x03, 0x04, 0x05, 0x06}, }, - expected: []interface{}{ - &Request{ + expected: []*Request{ + { SrcIP: net.IPv4(192, 168, 0, 3), SrcMAC: net.HardwareAddr{0x01, 0x02, 0x03, 0x04, 0x05, 0x06}, DstIP: net.IPv4(192, 168, 0, 1), @@ -775,24 +717,16 @@ func TestFileIPPortGenerator(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - reqgen := NewFileIPPortGenerator(func() (io.ReadCloser, error) { - return io.NopCloser(strings.NewReader(tt.input)), nil - }) - if tt.scanRange == nil { - tt.scanRange = &Range{} - } - pairs, err := reqgen.GenerateRequests(context.Background(), tt.scanRange) - if !assert.NoError(t, err) { - return - } - result := collectInterfaces(chanPairToGeneric(pairs)) - assert.Equal(t, tt.expected, result) - }() - waitDone(t, done) + reqgen := NewFileIPPortGenerator(func() (io.ReadCloser, error) { + return io.NopCloser(strings.NewReader(tt.input)), nil + }) + if tt.scanRange == nil { + tt.scanRange = &Range{} + } + pairs, err := reqgen.GenerateRequests(context.Background(), tt.scanRange) + require.NoError(t, err) + result := collectChannel(pairs) + require.Equal(t, tt.expected, result) }) } } @@ -813,20 +747,20 @@ func TestFileIPGenerator(t *testing.T) { tests := []struct { name string input string - expected []interface{} + expected []GeneratorResult[net.IP] }{ { name: "OneIP", input: `{"ip":"192.168.0.1"}`, - expected: []interface{}{ - WrapIP(net.IPv4(192, 168, 0, 1)), + expected: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(192, 168, 0, 1)), }, }, { name: "OneIPWithUnknownField", input: `{"ip":"192.168.0.1","abc":"field"}`, - expected: []interface{}{ - WrapIP(net.IPv4(192, 168, 0, 1)), + expected: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(192, 168, 0, 1)), }, }, { @@ -835,16 +769,16 @@ func TestFileIPGenerator(t *testing.T) { `{"ip":"192.168.0.1"}`, `{"ip":"192.168.0.2"}`, }, "\n"), - expected: []interface{}{ - WrapIP(net.IPv4(192, 168, 0, 1)), - WrapIP(net.IPv4(192, 168, 0, 2)), + expected: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(192, 168, 0, 1)), + generatedIP(net.IPv4(192, 168, 0, 2)), }, }, { name: "InvalidJSON", input: `{"ip":"192`, - expected: []interface{}{ - &ipError{error: ErrJSON}, + expected: []GeneratorResult[net.IP]{ + generationError[net.IP](ErrJSON), }, }, { @@ -853,9 +787,9 @@ func TestFileIPGenerator(t *testing.T) { `{"ip":"192.168.0.1","port":888}`, `{"ip":"192`, }, "\n"), - expected: []interface{}{ - WrapIP(net.IPv4(192, 168, 0, 1)), - &ipError{error: ErrJSON}, + expected: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(192, 168, 0, 1)), + generationError[net.IP](ErrJSON), }, }, { @@ -865,16 +799,27 @@ func TestFileIPGenerator(t *testing.T) { `{"ip":"192`, `{"ip":"192.168.0.3","port":888}`, }, "\n"), - expected: []interface{}{ - WrapIP(net.IPv4(192, 168, 0, 1)), - &ipError{error: ErrJSON}, + expected: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(192, 168, 0, 1)), + generationError[net.IP](ErrJSON), }, }, { name: "InvalidIP", input: `{"ip":"192.168.0.1111"}`, - expected: []interface{}{ - &ipError{error: ErrIP}, + expected: []GeneratorResult[net.IP]{ + generationError[net.IP](ErrIP), + }, + }, + { + name: "EmptyIPAfterValid", + input: strings.Join([]string{ + `{"ip":"192.168.0.1"}`, + `{}`, + }, "\n"), + expected: []GeneratorResult[net.IP]{ + generatedIP(net.IPv4(192, 168, 0, 1)), + generationError[net.IP](ErrIP), }, }, } @@ -883,25 +828,41 @@ func TestFileIPGenerator(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - ipgen := NewFileIPGenerator(func() (io.ReadCloser, error) { - return io.NopCloser(strings.NewReader(tt.input)), nil - }) - ips, err := ipgen.IPs(context.Background(), &Range{}) - if !assert.NoError(t, err) { - return - } - result := collectInterfaces(chanIPToGeneric(ips)) - assert.Equal(t, tt.expected, result) - }() - waitDone(t, done) + ipgen := NewFileIPGenerator(func() (io.ReadCloser, error) { + return io.NopCloser(strings.NewReader(tt.input)), nil + }) + ips, err := ipgen.IPs(context.Background(), &Range{}) + require.NoError(t, err) + result := collectChannel(ips) + require.Equal(t, tt.expected, result) }) } } +func TestFileIPGeneratorStopsAfterContextCancel(t *testing.T) { + t.Parallel() + + reader, writer := io.Pipe() + t.Cleanup(func() { + require.NoError(t, writer.Close()) + }) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + ips, err := NewFileIPGenerator(func() (io.ReadCloser, error) { + return reader, nil + }).IPs(ctx, &Range{}) + require.NoError(t, err) + _, err = writer.Write([]byte(`{"ip":"192.168.0.1"}` + "\n")) + require.NoError(t, err) + + select { + case _, ok := <-ips: + require.False(t, ok, "IP channel is not closed") + case <-time.After(100 * time.Millisecond): + require.FailNow(t, "file IP generator did not stop after context cancellation") + } +} + func TestLiveRequestGeneratorContextExit(t *testing.T) { t.Parallel() @@ -925,6 +886,38 @@ loop: } } +func TestLiveRequestGeneratorReportsRescanError(t *testing.T) { + t.Parallel() + + ctrl := gomock.NewController(t) + delegate := NewMockRequestGenerator(ctrl) + r := newScanRange() + initial := make(chan *Request) + close(initial) + rescanErr := errors.New("rescan error") + gomock.InOrder( + delegate.EXPECT().GenerateRequests(gomock.Any(), r).Return(initial, nil), + delegate.EXPECT().GenerateRequests(gomock.Any(), r).Return(nil, rescanErr), + ) + + rg := NewLiveRequestGenerator(delegate, time.Millisecond) + requests, err := rg.GenerateRequests(context.Background(), r) + require.NoError(t, err) + + select { + case request := <-requests: + require.ErrorIs(t, request.Err, rescanErr) + case <-time.After(waitTimeout): + require.FailNow(t, "rescan error was not reported") + } + select { + case _, ok := <-requests: + require.False(t, ok, "request channel is not closed") + case <-time.After(waitTimeout): + require.FailNow(t, "request channel was not closed") + } +} + func TestFilterIPRequestGenerator(t *testing.T) { t.Parallel() @@ -932,7 +925,7 @@ func TestFilterIPRequestGenerator(t *testing.T) { name string input []*Request filtered []bool - expected []interface{} + expected []*Request }{ { name: "EmptyFilter", @@ -940,7 +933,7 @@ func TestFilterIPRequestGenerator(t *testing.T) { newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())), newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), }, - expected: []interface{}{ + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())), newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), }, @@ -952,7 +945,7 @@ func TestFilterIPRequestGenerator(t *testing.T) { newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), }, filtered: []bool{true, false}, - expected: []interface{}{ + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), }, }, @@ -964,7 +957,7 @@ func TestFilterIPRequestGenerator(t *testing.T) { newScanRequest(withDstIP(net.IPv4(10, 0, 3, 3).To4())), }, filtered: []bool{false, true, false}, - expected: []interface{}{ + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())), newScanRequest(withDstIP(net.IPv4(10, 0, 3, 3).To4())), }, @@ -977,7 +970,7 @@ func TestFilterIPRequestGenerator(t *testing.T) { newScanRequest(withDstIP(net.IPv4(10, 0, 3, 3).To4())), }, filtered: []bool{true, false, true}, - expected: []interface{}{ + expected: []*Request{ newScanRequest(withDstIP(net.IPv4(10, 0, 2, 2).To4())), }, }, @@ -988,44 +981,36 @@ func TestFilterIPRequestGenerator(t *testing.T) { t.Run(tt.name, func(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) + ctrl := gomock.NewController(t) + delegate := NewMockRequestGenerator(ctrl) - ctrl := gomock.NewController(t) - delegate := NewMockRequestGenerator(ctrl) - - input := make(chan *Request, len(tt.input)) - for _, in := range tt.input { - input <- in - } - close(input) - r := newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}), - ) - delegate.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), r). - Return(input, nil) - - excludeIPs := NewMockIPContainer(ctrl) - var excludeFilters []gomock.Matcher - for i, filtered := range tt.filtered { - if filtered { - excludeIPs.EXPECT().Contains(tt.input[i].DstIP).Return(true, nil) - excludeFilters = append(excludeFilters, gomock.Not(gomock.Eq(tt.input[i].DstIP))) - } + input := make(chan *Request, len(tt.input)) + for _, in := range tt.input { + input <- in + } + close(input) + r := newScanRange( + withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}), + ) + delegate.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), r). + Return(input, nil) + + excludeIPs := NewMockIPContainer(ctrl) + var excludeFilters []gomock.Matcher + for i, filtered := range tt.filtered { + if filtered { + excludeIPs.EXPECT().Contains(tt.input[i].DstIP).Return(true, nil) + excludeFilters = append(excludeFilters, gomock.Not(gomock.Eq(tt.input[i].DstIP))) } - excludeIPs.EXPECT().Contains(gomock.All(excludeFilters...)).Return(false, nil).AnyTimes() + } + excludeIPs.EXPECT().Contains(gomock.All(excludeFilters...)).Return(false, nil).AnyTimes() - reqgen := NewFilterIPRequestGenerator(delegate, excludeIPs) - requests, err := reqgen.GenerateRequests(context.Background(), r) + reqgen := NewFilterIPRequestGenerator(delegate, excludeIPs) + requests, err := reqgen.GenerateRequests(context.Background(), r) - if !assert.NoError(t, err) { - return - } - result := collectInterfaces(chanPairToGeneric(requests)) - assert.Equal(t, tt.expected, result) - }() - waitDone(t, done) + require.NoError(t, err) + result := collectChannel(requests) + require.Equal(t, tt.expected, result) }) } } @@ -1033,61 +1018,47 @@ func TestFilterIPRequestGenerator(t *testing.T) { func TestFilterIPRequestGeneratorWithGeneratorError(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) + ctrl := gomock.NewController(t) + delegate := NewMockRequestGenerator(ctrl) - ctrl := gomock.NewController(t) - delegate := NewMockRequestGenerator(ctrl) + r := newScanRange( + withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}), + ) + delegate.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), r). + Return(nil, errors.New("generate error")) - r := newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}), - ) - delegate.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), r). - Return(nil, errors.New("generate error")) + excludeIPs := NewMockIPContainer(ctrl) + reqgen := NewFilterIPRequestGenerator(delegate, excludeIPs) + _, err := reqgen.GenerateRequests(context.Background(), r) - excludeIPs := NewMockIPContainer(ctrl) - reqgen := NewFilterIPRequestGenerator(delegate, excludeIPs) - _, err := reqgen.GenerateRequests(context.Background(), r) - - assert.Error(t, err) - }() - waitDone(t, done) + require.Error(t, err) } func TestFilterIPRequestGeneratorWithIPContainerError(t *testing.T) { t.Parallel() - done := make(chan interface{}) - go func() { - defer close(done) - - ctrl := gomock.NewController(t) - delegate := NewMockRequestGenerator(ctrl) + ctrl := gomock.NewController(t) + delegate := NewMockRequestGenerator(ctrl) - r := newScanRange( - withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}), - ) - input := make(chan *Request, 1) - input <- newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())) - close(input) - delegate.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), r). - Return(input, nil) + r := newScanRange( + withSubnet(&net.IPNet{IP: net.IPv4(10, 0, 0, 0), Mask: net.CIDRMask(8, 32)}), + ) + input := make(chan *Request, 1) + input <- newScanRequest(withDstIP(net.IPv4(10, 0, 1, 1).To4())) + close(input) + delegate.EXPECT().GenerateRequests(gomock.Not(gomock.Nil()), r). + Return(input, nil) - excludeIPs := NewMockIPContainer(ctrl) - excludeIPs.EXPECT().Contains(gomock.Any()).Return(false, errors.New("ip container error")) + excludeIPs := NewMockIPContainer(ctrl) + excludeIPs.EXPECT().Contains(gomock.Any()).Return(false, errors.New("ip container error")) - reqgen := NewFilterIPRequestGenerator(delegate, excludeIPs) - requests, err := reqgen.GenerateRequests(context.Background(), r) + reqgen := NewFilterIPRequestGenerator(delegate, excludeIPs) + requests, err := reqgen.GenerateRequests(context.Background(), r) - if !assert.NoError(t, err) { - return - } - result := collectInterfaces(chanPairToGeneric(requests)) - assert.Equal(t, []interface{}{ - newScanRequest( - withDstIP(net.IPv4(10, 0, 1, 1).To4()), - withError(errors.New("ip container error")))}, result) - }() - waitDone(t, done) + require.NoError(t, err) + result := collectChannel(requests) + require.Equal(t, []*Request{ + newScanRequest( + withDstIP(net.IPv4(10, 0, 1, 1).To4()), + withError(errors.New("ip container error")))}, result) } diff --git a/pkg/scan/tcp/tcp_test.go b/pkg/scan/tcp/tcp_test.go index 96e5b34..0a9a097 100644 --- a/pkg/scan/tcp/tcp_test.go +++ b/pkg/scan/tcp/tcp_test.go @@ -479,9 +479,15 @@ func TestAllFlags(t *testing.T) { } } -type mockIPGeneratorFunc func(ctx context.Context, r *scan.Range) (<-chan scan.IPGetter, error) - -func (f mockIPGeneratorFunc) IPs(ctx context.Context, r *scan.Range) (<-chan scan.IPGetter, error) { +type mockIPGeneratorFunc func( + ctx context.Context, + r *scan.Range, +) (<-chan scan.GeneratorResult[net.IP], error) + +func (f mockIPGeneratorFunc) IPs( + ctx context.Context, + r *scan.Range, +) (<-chan scan.GeneratorResult[net.IP], error) { return f(ctx, r) } @@ -501,15 +507,18 @@ func BenchmarkTCPScanEngine(b *testing.B) { defer cancel() dstIP := net.IPv4(192, 168, 0, 3).To4() - ipgen := mockIPGeneratorFunc(func(ctx context.Context, r *scan.Range) (<-chan scan.IPGetter, error) { - out := make(chan scan.IPGetter, 100) + ipgen := mockIPGeneratorFunc(func( + ctx context.Context, + _ *scan.Range, + ) (<-chan scan.GeneratorResult[net.IP], error) { + out := make(chan scan.GeneratorResult[net.IP], 100) go func() { defer close(out) for i := 0; i < b.N; i++ { select { case <-ctx.Done(): return - case out <- scan.WrapIP(dstIP): + case out <- scan.GeneratorResult[net.IP]{Value: dstIP}: } } }() diff --git a/pkg/scan/utils_test.go b/pkg/scan/utils_test.go index 18890bd..062ab96 100644 --- a/pkg/scan/utils_test.go +++ b/pkg/scan/utils_test.go @@ -9,9 +9,9 @@ import ( const waitTimeout = 3 * time.Second -func chanToSlice(t *testing.T, in <-chan interface{}, expectedLen int) []interface{} { +func chanToSlice[T any](t *testing.T, in <-chan T, expectedLen int) []T { t.Helper() - result := []interface{}{} + result := []T{} loop: for { select { @@ -24,23 +24,12 @@ loop: } result = append(result, data) case <-time.After(waitTimeout): - t.Fatal("read timeout") + require.FailNow(t, "read timeout") } } return result } -func chanErrToGeneric(in <-chan error) <-chan interface{} { - out := make(chan interface{}, cap(in)) - go func() { - defer close(out) - for i := range in { - out <- i - } - }() - return out -} - func waitDone(t *testing.T, done <-chan interface{}) { t.Helper() select {