коммит произвёл
GitHub
родитель
87a719b4b4
Коммит
3c625743e5
@@ -4,6 +4,7 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"io/ioutil"
|
||||
@@ -112,3 +113,23 @@ func CopyDir(src string, dst string) (err error) {
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
var SizeLimitExceeded = errors.New("Size limit exceeded")
|
||||
|
||||
type LimitedReaderWithError struct {
|
||||
limitedReader *io.LimitedReader
|
||||
}
|
||||
|
||||
func NewLimitedReaderWithError(reader io.Reader, maxBytes int64) *LimitedReaderWithError {
|
||||
return &LimitedReaderWithError{
|
||||
limitedReader: &io.LimitedReader{R: reader, N: maxBytes + 1},
|
||||
}
|
||||
}
|
||||
|
||||
func (l *LimitedReaderWithError) Read(p []byte) (int, error) {
|
||||
n, err := l.limitedReader.Read(p)
|
||||
if l.limitedReader.N <= 0 && err == io.EOF {
|
||||
return n, SizeLimitExceeded
|
||||
}
|
||||
return n, err
|
||||
}
|
||||
|
||||
@@ -4,6 +4,9 @@
|
||||
package utils
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/rand"
|
||||
"io"
|
||||
"io/ioutil"
|
||||
"os"
|
||||
"path/filepath"
|
||||
@@ -62,3 +65,65 @@ func TestCopyDir(t *testing.T) {
|
||||
err = CopyDir(srcDir, dstDir)
|
||||
assert.Error(t, err)
|
||||
}
|
||||
func TestLimitedReaderWithError(t *testing.T) {
|
||||
t.Run("read less than max size", func(t *testing.T) {
|
||||
maxBytes := 10
|
||||
randomBytes := make([]byte, maxBytes)
|
||||
n, err := rand.Read(randomBytes)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, n, maxBytes)
|
||||
|
||||
lr := NewLimitedReaderWithError(bytes.NewReader(randomBytes), int64(maxBytes))
|
||||
smallerBuf := make([]byte, maxBytes-3)
|
||||
_, err = io.ReadFull(lr, smallerBuf)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
|
||||
t.Run("read equal to max size", func(t *testing.T) {
|
||||
maxBytes := 10
|
||||
randomBytes := make([]byte, maxBytes)
|
||||
n, err := rand.Read(randomBytes)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, n, maxBytes)
|
||||
|
||||
lr := NewLimitedReaderWithError(bytes.NewReader(randomBytes), int64(maxBytes))
|
||||
buf := make([]byte, maxBytes)
|
||||
_, err = io.ReadFull(lr, buf)
|
||||
require.Truef(t, err == nil || err == io.EOF, "err must be nil or %v, got %v", io.EOF, err)
|
||||
})
|
||||
|
||||
t.Run("single read, larger than max size", func(t *testing.T) {
|
||||
maxBytes := 5
|
||||
moreThanMaxBytes := maxBytes + 10
|
||||
randomBytes := make([]byte, moreThanMaxBytes)
|
||||
n, err := rand.Read(randomBytes)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, moreThanMaxBytes, n)
|
||||
|
||||
lr := NewLimitedReaderWithError(bytes.NewReader(randomBytes), int64(maxBytes))
|
||||
buf := make([]byte, moreThanMaxBytes)
|
||||
_, err = io.ReadFull(lr, buf)
|
||||
require.Error(t, err)
|
||||
require.Equal(t, SizeLimitExceeded, err)
|
||||
})
|
||||
|
||||
t.Run("multiple small reads, total larger than max size", func(t *testing.T) {
|
||||
maxBytes := 10
|
||||
lessThanMaxBytes := maxBytes - 4
|
||||
randomBytesLen := maxBytes * 2
|
||||
randomBytes := make([]byte, randomBytesLen)
|
||||
n, err := rand.Read(randomBytes)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, randomBytesLen, n)
|
||||
|
||||
lr := NewLimitedReaderWithError(bytes.NewReader(randomBytes), int64(maxBytes))
|
||||
buf := make([]byte, lessThanMaxBytes)
|
||||
_, err = io.ReadFull(lr, buf)
|
||||
require.NoError(t, err)
|
||||
|
||||
// lets do it again
|
||||
_, err = io.ReadFull(lr, buf)
|
||||
require.Error(t, err)
|
||||
require.Equal(t, SizeLimitExceeded, err)
|
||||
})
|
||||
}
|
||||
|
||||
Ссылка в новой задаче
Block a user