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
7 changes: 5 additions & 2 deletions common/httpx/filter.go
Original file line number Diff line number Diff line change
Expand Up @@ -53,8 +53,11 @@ type FilterCustom struct {
func (f FilterCustom) Filter(response *Response) (bool, error) {
for _, callback := range f.CallBacks {
ok, err := callback(response)
if ok && err == nil {
return true, err
if err != nil {
return false, err
}
if ok {
return true, nil
}
}

Expand Down
73 changes: 73 additions & 0 deletions common/httpx/filter_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
package httpx

import (
"errors"
"testing"

"github.com/stretchr/testify/require"
)

func TestFilterCustomErrorPropagation(t *testing.T) {
t.Run("error from callback is returned, not swallowed", func(t *testing.T) {
expectedErr := errors.New("callback failure")
callback := func(response *Response) (bool, error) {
return true, expectedErr
}
filter := FilterCustom{CallBacks: []CustomCallback{callback}}
ok, err := filter.Filter(&Response{})
require.False(t, ok, "ok should be false when callback returns an error")
require.ErrorIs(t, err, expectedErr, "error from callback should be propagated")
})

t.Run("error from callback with ok=false is returned", func(t *testing.T) {
expectedErr := errors.New("callback failure")
callback := func(response *Response) (bool, error) {
return false, expectedErr
}
filter := FilterCustom{CallBacks: []CustomCallback{callback}}
ok, err := filter.Filter(&Response{})
require.False(t, ok)
require.ErrorIs(t, err, expectedErr)
})

t.Run("first matching callback without error returns true", func(t *testing.T) {
callbacks := []CustomCallback{
func(response *Response) (bool, error) { return false, nil },
func(response *Response) (bool, error) { return true, nil },
func(response *Response) (bool, error) { return true, nil },
}
filter := FilterCustom{CallBacks: callbacks}
ok, err := filter.Filter(&Response{})
require.True(t, ok)
require.NoError(t, err)
})

t.Run("error stops remaining callbacks", func(t *testing.T) {
called := 0
filter := FilterCustom{CallBacks: []CustomCallback{
func(*Response) (bool, error) {
called++
return false, errors.New("fail")
},
func(*Response) (bool, error) {
called++
return true, nil
},
}}
ok, err := filter.Filter(&Response{})
require.False(t, ok)
require.Error(t, err)
require.Equal(t, 1, called)
})

t.Run("no callbacks match returns false with nil error", func(t *testing.T) {
callbacks := []CustomCallback{
func(response *Response) (bool, error) { return false, nil },
func(response *Response) (bool, error) { return false, nil },
}
filter := FilterCustom{CallBacks: callbacks}
ok, err := filter.Filter(&Response{})
require.False(t, ok)
require.NoError(t, err)
})
}
1 change: 1 addition & 0 deletions common/httpx/pipeline.go
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ func (h *HTTPX) SupportPipeline(protocol, method, host string, port int) bool {
if err != nil {
return false
}
defer func() { _ = conn.Close() }()
// send some probes
nprobes := 10
for i := 0; i < nprobes; i++ {
Expand Down
51 changes: 51 additions & 0 deletions common/httpx/pipeline_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,51 @@
package httpx

import (
"io"
"net"
"strconv"
"testing"
"time"

"github.com/stretchr/testify/require"
)

func TestSupportPipelineClosesConn(t *testing.T) {
ln, err := net.Listen("tcp", "127.0.0.1:0")
require.NoError(t, err)
defer func() { _ = ln.Close() }()

_, portStr, err := net.SplitHostPort(ln.Addr().String())
require.NoError(t, err)
port, err := strconv.Atoi(portStr)
require.NoError(t, err)

closed := make(chan struct{})
go func() {
conn, err := ln.Accept()
if err != nil {
return
}
defer func() { _ = conn.Close() }()
buf := make([]byte, 64*1024)
_ = conn.SetReadDeadline(time.Now().Add(2 * time.Second))
_, _ = conn.Read(buf)
_, _ = conn.Write([]byte("HTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\nHTTP/1.1 200 OK\r\nContent-Length: 0\r\n\r\n"))
_, _ = io.Copy(io.Discard, conn)
close(closed)
}()

h := &HTTPX{}
_ = h.SupportPipeline("http", "GET", "127.0.0.1", port)

select {
case <-closed:
case <-time.After(3 * time.Second):
t.Fatal("pipeline probe connection was not closed")
}
}

func TestSupportPipelineDialError(t *testing.T) {
h := &HTTPX{}
require.False(t, h.SupportPipeline("http", "GET", "127.0.0.1", 1))
}
29 changes: 19 additions & 10 deletions internal/pdcp/writer.go
Original file line number Diff line number Diff line change
Expand Up @@ -131,6 +131,7 @@ func (u *UploadWriter) autoCommit(ctx context.Context) {
// temporary buffer to store the results
buff := &bytes.Buffer{}
ticker := time.NewTicker(flushTimer)
defer ticker.Stop()

for {
select {
Expand Down Expand Up @@ -169,15 +170,13 @@ func (u *UploadWriter) autoCommit(ctx context.Context) {
}
u.counter.Add(1)
line := conversion.String(lineBytes)
if buff.Len()+len(line) > MaxChunkSize {
// flush existing buffer
if err := u.uploadChunk(buff); err != nil {
appendResultLine(buff, line, MaxChunkSize, func(b *bytes.Buffer) error {
if err := u.uploadChunk(b); err != nil {
gologger.Error().Msgf("Failed to upload asset results on cloud: %v", err)
return err
}
} else {
buff.WriteString(line)
buff.WriteString("\n")
}
return nil
})
}
}
}
Expand Down Expand Up @@ -259,12 +258,22 @@ func (u *UploadWriter) getRequest(bin []byte) (*retryablehttp.Request, error) {
return req, nil
}

// appendResultLine writes line and its trailing newline to buff, flushing
// existing data first when the next write would exceed max. An empty buffer is
// never flushed, so a single oversized line is retained instead of dropped or
// uploaded as an empty chunk.
func appendResultLine(buff *bytes.Buffer, line string, max int, flush func(*bytes.Buffer) error) {
if buff.Len() > 0 && buff.Len()+len(line)+len("\n") > max {
_ = flush(buff)
}
buff.WriteString(line)
buff.WriteString("\n")
}

// Close closes the upload writer
func (u *UploadWriter) Close() {
if !u.closed.Load() {
// protect to avoid channel closed twice error
if u.closed.CompareAndSwap(false, true) {
close(u.data)
u.closed.Store(true)
}
<-u.done
}
131 changes: 131 additions & 0 deletions internal/pdcp/writer_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,131 @@
package pdcp

import (
"bytes"
"errors"
"sync"
"testing"
"time"

"github.com/projectdiscovery/httpx/runner"
"github.com/stretchr/testify/require"
)

func TestAppendResultLine(t *testing.T) {
t.Run("keeps lines under the limit without flushing", func(t *testing.T) {
buff := &bytes.Buffer{}
flushed := 0
appendResultLine(buff, "ab", 10, func(*bytes.Buffer) error {
flushed++
return nil
})
appendResultLine(buff, "cd", 10, func(*bytes.Buffer) error {
flushed++
return nil
})
require.Equal(t, 0, flushed)
require.Equal(t, "ab\ncd\n", buff.String())
})

t.Run("flushes existing data and keeps the overflowing line", func(t *testing.T) {
buff := &bytes.Buffer{}
flush := func(b *bytes.Buffer) error {
require.Equal(t, "aaaa\n", b.String())
b.Reset()
return nil
}
appendResultLine(buff, "aaaa", 6, flush)
appendResultLine(buff, "bbbb", 6, flush)
require.Equal(t, "bbbb\n", buff.String())
})

t.Run("does not flush an empty buffer for an oversized line", func(t *testing.T) {
buff := &bytes.Buffer{}
flushed := 0
appendResultLine(buff, "toolong", 4, func(*bytes.Buffer) error {
flushed++
return nil
})
require.Equal(t, 0, flushed)
require.Equal(t, "toolong\n", buff.String())
})

t.Run("newline counts towards the limit", func(t *testing.T) {
buff := &bytes.Buffer{}
const max = 6
flush := func(b *bytes.Buffer) error {
b.Reset()
return nil
}
// "abc\n" is 4 bytes, appending "de\n" would reach 7 without counting
// the newline in the check.
appendResultLine(buff, "abc", max, flush)
appendResultLine(buff, "de", max, flush)
require.LessOrEqual(t, buff.Len(), max)
require.Equal(t, "de\n", buff.String())
})

t.Run("still appends the current line when flush fails", func(t *testing.T) {
buff := bytes.NewBufferString("old\n")
appendResultLine(buff, "new", 4, func(*bytes.Buffer) error {
return errors.New("upload failed")
})
require.Equal(t, "old\nnew\n", buff.String())
})
}

func TestUploadWriterCloseWaits(t *testing.T) {
u := &UploadWriter{
done: make(chan struct{}, 1),
data: make(chan runner.Result, 8),
}

started := make(chan struct{})
release := make(chan struct{})
go func() {
for range u.data {
}
close(started)
<-release
u.done <- struct{}{}
close(u.done)
}()

firstDone := make(chan struct{})
go func() {
u.Close()
close(firstDone)
}()

select {
case <-started:
case <-time.After(2 * time.Second):
t.Fatal("Close did not close the data channel")
}

secondDone := make(chan struct{})
go func() {
u.Close()
close(secondDone)
}()

select {
case <-secondDone:
t.Fatal("second Close returned before autoCommit finished")
case <-time.After(50 * time.Millisecond):
}

close(release)

var wg sync.WaitGroup
wg.Add(2)
go func() { defer wg.Done(); <-firstDone }()
go func() { defer wg.Done(); <-secondDone }()
done := make(chan struct{})
go func() { wg.Wait(); close(done) }()
select {
case <-done:
case <-time.After(2 * time.Second):
t.Fatal("Close hung")
}
}
Loading