Skip to content

Commit 35384b4

Browse files
authored
feat: add bounded read helpers (aquasecurity#10974)
1 parent 3c6a1a2 commit 35384b4

4 files changed

Lines changed: 266 additions & 0 deletions

File tree

pkg/x/io/max_bytes_reader.go

Lines changed: 71 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,71 @@
1+
package io
2+
3+
import (
4+
"errors"
5+
"fmt"
6+
"io"
7+
)
8+
9+
// ErrLimitExceeded identifies reads that exceed a configured byte limit.
10+
var ErrLimitExceeded = errors.New("io: read limit exceeded")
11+
12+
// MaxBytesError is returned by MaxBytesReader when its read limit is exceeded.
13+
type MaxBytesError struct {
14+
Limit int64
15+
}
16+
17+
func (e *MaxBytesError) Error() string {
18+
return fmt.Sprintf("%s: %d-byte limit", ErrLimitExceeded, e.Limit)
19+
}
20+
21+
func (*MaxBytesError) Is(target error) bool {
22+
return target == ErrLimitExceeded
23+
}
24+
25+
// MaxBytesReader returns a reader that exposes at most n bytes from r. It
26+
// returns a MaxBytesError when a read proves that r contains more than n bytes.
27+
// Detecting overflow consumes one byte beyond the limit from r.
28+
// A negative limit is treated as zero.
29+
func MaxBytesReader(r io.Reader, n int64) io.Reader {
30+
if n < 0 {
31+
n = 0
32+
}
33+
return &maxBytesReader{
34+
r: r,
35+
limit: n,
36+
remaining: n,
37+
}
38+
}
39+
40+
type maxBytesReader struct {
41+
r io.Reader
42+
limit int64
43+
remaining int64
44+
err error
45+
}
46+
47+
func (r *maxBytesReader) Read(p []byte) (int, error) {
48+
if r.err != nil {
49+
return 0, r.err
50+
}
51+
if len(p) == 0 {
52+
return 0, nil
53+
}
54+
if int64(len(p))-1 > r.remaining {
55+
p = p[:r.remaining+1]
56+
}
57+
58+
n, err := r.r.Read(p)
59+
if int64(n) <= r.remaining {
60+
r.remaining -= int64(n)
61+
r.err = err
62+
return n, err
63+
}
64+
65+
// The overflow branch guarantees r.remaining < int64(n), so converting
66+
// r.remaining to int is safe.
67+
n = int(r.remaining)
68+
r.remaining = 0
69+
r.err = &MaxBytesError{Limit: r.limit}
70+
return n, r.err
71+
}

pkg/x/io/max_bytes_reader_test.go

Lines changed: 141 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,141 @@
1+
package io_test
2+
3+
import (
4+
"errors"
5+
"fmt"
6+
"io"
7+
"math"
8+
"strings"
9+
"testing"
10+
11+
"github.com/stretchr/testify/assert"
12+
"github.com/stretchr/testify/require"
13+
14+
xio "github.com/aquasecurity/trivy/pkg/x/io"
15+
)
16+
17+
type readerFunc func([]byte) (int, error)
18+
19+
func (f readerFunc) Read(p []byte) (int, error) {
20+
return f(p)
21+
}
22+
23+
func TestMaxBytesReader(t *testing.T) {
24+
tests := []struct {
25+
name string
26+
input string
27+
limit int64
28+
want string
29+
wantError bool
30+
wantLimit int64
31+
}{
32+
{name: "empty input at zero limit", limit: 0},
33+
{name: "input below limit", input: "abc", limit: 4, want: "abc"},
34+
{name: "input exactly at limit", input: "abc", limit: 3, want: "abc"},
35+
{name: "input one byte over limit", input: "abcd", limit: 3, want: "abc", wantError: true, wantLimit: 3},
36+
{name: "non-empty input at zero limit", input: "a", wantError: true, wantLimit: 0},
37+
{name: "negative limit is zero", input: "a", limit: -1, wantError: true, wantLimit: 0},
38+
{name: "maximum limit", input: "abc", limit: math.MaxInt64, want: "abc"},
39+
}
40+
41+
for _, tt := range tests {
42+
t.Run(tt.name, func(t *testing.T) {
43+
got, err := io.ReadAll(xio.MaxBytesReader(strings.NewReader(tt.input), tt.limit))
44+
assert.Equal(t, tt.want, string(got))
45+
if !tt.wantError {
46+
require.NoError(t, err)
47+
return
48+
}
49+
50+
require.ErrorIs(t, err, xio.ErrLimitExceeded)
51+
assert.Equal(t, fmt.Sprintf("io: read limit exceeded: %d-byte limit", tt.wantLimit), err.Error())
52+
53+
wrapped := fmt.Errorf("read input: %w", err)
54+
require.ErrorIs(t, wrapped, xio.ErrLimitExceeded)
55+
var maxErr *xio.MaxBytesError
56+
require.ErrorAs(t, wrapped, &maxErr)
57+
assert.Equal(t, tt.wantLimit, maxErr.Limit)
58+
})
59+
}
60+
}
61+
62+
func TestMaxBytesReaderMultipleReads(t *testing.T) {
63+
r := xio.MaxBytesReader(strings.NewReader("abcd"), 3)
64+
buf := make([]byte, 2)
65+
66+
n, err := r.Read(buf)
67+
require.NoError(t, err)
68+
assert.Equal(t, 2, n)
69+
assert.Equal(t, "ab", string(buf[:n]))
70+
71+
n, err = r.Read(buf)
72+
require.ErrorIs(t, err, xio.ErrLimitExceeded)
73+
assert.Equal(t, 1, n)
74+
assert.Equal(t, "c", string(buf[:n]))
75+
limitErr := err
76+
77+
n, err = r.Read(buf)
78+
assert.Equal(t, 0, n)
79+
assert.Same(t, limitErr, err)
80+
}
81+
82+
func TestMaxBytesReaderUnderlyingError(t *testing.T) {
83+
sourceErr := errors.New("source error")
84+
reads := 0
85+
r := xio.MaxBytesReader(readerFunc(func(p []byte) (int, error) {
86+
reads++
87+
return copy(p, "ab"), sourceErr
88+
}), 3)
89+
buf := make([]byte, 4)
90+
91+
n, err := r.Read(buf)
92+
assert.Equal(t, 2, n)
93+
assert.Equal(t, "ab", string(buf[:n]))
94+
assert.Same(t, sourceErr, err)
95+
96+
n, err = r.Read(buf)
97+
assert.Equal(t, 0, n)
98+
assert.Same(t, sourceErr, err)
99+
assert.Equal(t, 1, reads)
100+
}
101+
102+
func TestMaxBytesReaderOverflowTakesPrecedence(t *testing.T) {
103+
sourceErr := errors.New("source error")
104+
r := xio.MaxBytesReader(readerFunc(func(p []byte) (int, error) {
105+
return copy(p, "abcd"), sourceErr
106+
}), 3)
107+
buf := make([]byte, 8)
108+
109+
n, err := r.Read(buf)
110+
assert.Equal(t, 3, n)
111+
assert.Equal(t, "abc", string(buf[:n]))
112+
require.ErrorIs(t, err, xio.ErrLimitExceeded)
113+
assert.NotErrorIs(t, err, sourceErr)
114+
}
115+
116+
func TestMaxBytesReaderReadBound(t *testing.T) {
117+
source := xio.NewCountingReader(strings.NewReader("abcdef"))
118+
r := xio.MaxBytesReader(source, 3)
119+
120+
got, err := io.ReadAll(r)
121+
assert.Equal(t, "abc", string(got))
122+
require.ErrorIs(t, err, xio.ErrLimitExceeded)
123+
assert.Equal(t, int64(4), source.BytesRead())
124+
limitErr := err
125+
126+
buf := make([]byte, 1)
127+
n, err := r.Read(buf)
128+
assert.Equal(t, 0, n)
129+
assert.Same(t, limitErr, err)
130+
assert.Equal(t, int64(4), source.BytesRead())
131+
}
132+
133+
func TestMaxBytesReaderZeroLengthRead(t *testing.T) {
134+
source := xio.NewCountingReader(strings.NewReader("a"))
135+
r := xio.MaxBytesReader(source, -1)
136+
137+
n, err := r.Read(nil)
138+
require.NoError(t, err)
139+
assert.Equal(t, 0, n)
140+
assert.Equal(t, int64(0), source.BytesRead())
141+
}

pkg/x/io/read_all.go

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,10 @@
1+
package io
2+
3+
import "io"
4+
5+
// ReadAllWithLimit reads from r until EOF or until the read exceeds n bytes.
6+
// On overflow, it returns the first n bytes and a MaxBytesError.
7+
// A negative limit is treated as zero.
8+
func ReadAllWithLimit(r io.Reader, n int64) ([]byte, error) {
9+
return io.ReadAll(MaxBytesReader(r, n))
10+
}

pkg/x/io/read_all_test.go

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
package io_test
2+
3+
import (
4+
"strings"
5+
"testing"
6+
7+
"github.com/stretchr/testify/assert"
8+
"github.com/stretchr/testify/require"
9+
10+
"github.com/aquasecurity/trivy/pkg/x/io"
11+
)
12+
13+
func TestReadAllWithLimit(t *testing.T) {
14+
tests := []struct {
15+
name string
16+
input string
17+
limit int64
18+
want string
19+
wantError bool
20+
wantLimit int64
21+
}{
22+
{name: "empty input at zero limit", limit: 0},
23+
{name: "input below limit", input: "abc", limit: 4, want: "abc"},
24+
{name: "input exactly at limit", input: "abc", limit: 3, want: "abc"},
25+
{name: "input over limit", input: "abcd", limit: 3, want: "abc", wantError: true, wantLimit: 3},
26+
{name: "negative limit", input: "a", limit: -1, wantError: true, wantLimit: 0},
27+
}
28+
29+
for _, tt := range tests {
30+
t.Run(tt.name, func(t *testing.T) {
31+
got, err := io.ReadAllWithLimit(strings.NewReader(tt.input), tt.limit)
32+
assert.Equal(t, tt.want, string(got))
33+
if !tt.wantError {
34+
require.NoError(t, err)
35+
return
36+
}
37+
38+
require.ErrorIs(t, err, io.ErrLimitExceeded)
39+
var maxErr *io.MaxBytesError
40+
require.ErrorAs(t, err, &maxErr)
41+
assert.Equal(t, tt.wantLimit, maxErr.Limit)
42+
})
43+
}
44+
}

0 commit comments

Comments
 (0)