diff --git a/common/httpx/filter.go b/common/httpx/filter.go index e553abcbb..44ac72ef0 100644 --- a/common/httpx/filter.go +++ b/common/httpx/filter.go @@ -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 } } diff --git a/common/httpx/filter_test.go b/common/httpx/filter_test.go new file mode 100644 index 000000000..8c96e43a4 --- /dev/null +++ b/common/httpx/filter_test.go @@ -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) + }) +} diff --git a/common/httpx/pipeline.go b/common/httpx/pipeline.go index b6b7b7817..e1f4b0188 100644 --- a/common/httpx/pipeline.go +++ b/common/httpx/pipeline.go @@ -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++ { diff --git a/common/httpx/pipeline_test.go b/common/httpx/pipeline_test.go new file mode 100644 index 000000000..01eb683ff --- /dev/null +++ b/common/httpx/pipeline_test.go @@ -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)) +} diff --git a/internal/pdcp/writer.go b/internal/pdcp/writer.go index fbad91abd..d397b06d6 100644 --- a/internal/pdcp/writer.go +++ b/internal/pdcp/writer.go @@ -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 { @@ -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 + }) } } } @@ -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 } diff --git a/internal/pdcp/writer_test.go b/internal/pdcp/writer_test.go new file mode 100644 index 000000000..a877eac6d --- /dev/null +++ b/internal/pdcp/writer_test.go @@ -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") + } +}