Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 0 additions & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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)

</div>
Expand Down
2 changes: 1 addition & 1 deletion command/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)))
Expand Down
4 changes: 2 additions & 2 deletions pkg/packet/mock_sender_test.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

29 changes: 14 additions & 15 deletions pkg/packet/receiver_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
)
Expand Down Expand Up @@ -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")
})
}
}
Expand All @@ -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) {
Expand All @@ -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) {
Expand Down Expand Up @@ -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")
})
}
}
Expand All @@ -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) {
Expand All @@ -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")
}
47 changes: 35 additions & 12 deletions pkg/packet/sender.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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() {
Expand All @@ -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
}
}
82 changes: 58 additions & 24 deletions pkg/packet/sender_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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) {
Expand All @@ -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) {
Expand Down Expand Up @@ -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) {
Expand All @@ -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) {
Expand All @@ -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) {
Expand All @@ -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")
}
}
17 changes: 3 additions & 14 deletions pkg/packet/utils_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand All @@ -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
}
45 changes: 45 additions & 0 deletions pkg/scan/channels.go
Original file line number Diff line number Diff line change
@@ -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
}
Loading