Skip to content
Open
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
179 changes: 179 additions & 0 deletions internal/stream/readat_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,179 @@
package stream_test

import (
"bytes"
"context"
"io"
"math/rand"
"sync/atomic"
"testing"

"github.com/OpenListTeam/OpenList/v4/internal/model"
"github.com/OpenListTeam/OpenList/v4/internal/stream"
"github.com/OpenListTeam/OpenList/v4/pkg/http_range"
)

// maxReuseGap mirrors the internal continuation-reuse window (4*utils.MB).
const maxReuseGap = 4 * 1024 * 1024

// newMockSeekableStream builds a SeekableStream whose range reads are served
// from data, counting every upstream range request in gets.
func newMockSeekableStream(t *testing.T, data []byte, gets *atomic.Int64) *stream.SeekableStream {
t.Helper()
rr := stream.RangeReaderFunc(func(ctx context.Context, r http_range.Range) (io.ReadCloser, error) {
gets.Add(1)
if r.Length < 0 || r.Start+r.Length > int64(len(data)) {
r.Length = int64(len(data)) - r.Start
}
return io.NopCloser(io.NewSectionReader(bytes.NewReader(data), r.Start, r.Length)), nil
})
ss, err := stream.NewSeekableStream(&stream.FileStream{
Obj: &model.Object{Size: int64(len(data))},
Ctx: context.Background(),
}, &model.Link{
RangeReader: rr,
ContentLength: int64(len(data)),
})
if err != nil {
t.Fatalf("NewSeekableStream() error = %v", err)
}
return ss
}

// readAtFull reads len(p) bytes at off and fails the test on mismatch.
func readAtFull(t *testing.T, ra io.ReaderAt, data []byte, off int64, p []byte) {
t.Helper()
n, err := ra.ReadAt(p, off)
if err != nil {
t.Fatalf("ReadAt(off=%d) error = %v", off, err)
}
if !bytes.Equal(p, data[off:off+int64(n)]) {
t.Fatalf("ReadAt(off=%d) content mismatch", off)
}
}

func randomData(size int) []byte {
data := make([]byte, size)
x := uint64(42)
for i := range data {
x = x*6364136223846793005 + 1
data[i] = byte(x >> 33)
}
return data
}

// Sequential reads must reuse a single upstream range request.
func TestReadAtSeekerSequentialReuse(t *testing.T) {
data := randomData(16 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
buf := make([]byte, 128*1024)
for off := 0; off < len(data); off += len(buf) {
readAtFull(t, ra, data, int64(off), buf)
}
if n := gets.Load(); n != 1 {
t.Fatalf("sequential read issued %d range requests, want 1", n)
}
}

// A read landing up to maxReuseGap bytes past a parked reader must be served
// by advancing that reader, without a new range request.
func TestReadAtSeekerSkipsAheadWithinWindow(t *testing.T) {
data := randomData(16 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
// Park a continuation reader right after reading the first 2 MiB.
chunk := make([]byte, 256*1024)
for off := 0; off < 2*1024*1024; off += len(chunk) {
readAtFull(t, ra, data, int64(off), chunk)
}
skip := 512 * 1024
off := int64(2*1024*1024 + skip)
readAtFull(t, ra, data, off, chunk)
if n := gets.Load(); n != 1 {
t.Fatalf("window skip issued %d range requests, want 1", n)
}
// A second skip deeper inside the window must also be free.
off = int64(4*1024*1024) - 128*1024
readAtFull(t, ra, data, off, chunk)
if n := gets.Load(); n != 1 {
t.Fatalf("second window skip issued %d range requests, want 1", n)
}
}

// A forward jump beyond the reuse window must open a new range request but
// keep the parked reader available for later window hits.
func TestReadAtSeekerFarJumpOpensNewRequest(t *testing.T) {
data := randomData(16 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
buf := make([]byte, 256*1024)
for off := 0; off < 2*1024*1024; off += len(buf) {
readAtFull(t, ra, data, int64(off), buf)
}
// 2 MiB -> 10 MiB is beyond the 4 MiB reuse window.
off := int64(10 * 1024 * 1024)
readAtFull(t, ra, data, off, buf)
if n := gets.Load(); n != 2 {
t.Fatalf("far jump issued %d range requests, want 2", n)
}
// Back within the window of the 10 MiB chain: free reuse again.
readAtFull(t, ra, data, off+maxReuseGap, buf)
if n := gets.Load(); n != 2 {
t.Fatalf("jump inside new window issued %d range requests, want 2", n)
}
}

// Backward reads can never reuse a parked continuation and must open a new
// range request.
func TestReadAtSeekerBackwardJumpOpensNewRequest(t *testing.T) {
data := randomData(8 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
buf := make([]byte, 256*1024)
for off := 0; off < 2*1024*1024; off += len(buf) {
readAtFull(t, ra, data, int64(off), buf)
}
readAtFull(t, ra, data, int64(1024*1024), buf)
if n := gets.Load(); n != 2 {
t.Fatalf("backward jump issued %d range requests, want 2", n)
}
}

// Random reads must return correct data and keep upstream requests bounded:
// each read is either a window hit or a fresh request, never more than one.
func TestReadAtSeekerRandomReads(t *testing.T) {
data := randomData(32 * 1024 * 1024)
var gets atomic.Int64
ss := newMockSeekableStream(t, data, &gets)
ra, err := stream.NewReadAtSeeker(ss, 0, true)
if err != nil {
t.Fatalf("NewReadAtSeeker() error = %v", err)
}
const chunk = 8 * 1024
buf := make([]byte, chunk)
r := rand.New(rand.NewSource(7))
for i := 0; i < 200; i++ {
off := r.Int63n(int64(len(data)) - chunk)
readAtFull(t, ra, data, off, buf)
}
if n := gets.Load(); n > 200 {
t.Fatalf("random reads issued %d range requests, want <= 200", n)
}
}
105 changes: 71 additions & 34 deletions internal/stream/stream.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"io"
"math"
"os"
"sort"
"sync"

"github.com/OpenListTeam/OpenList/v4/internal/conf"
Expand Down Expand Up @@ -358,10 +359,72 @@ func (r *ReaderUpdatingProgress) Close() error {
type RangeReadReadAtSeeker struct {
ss *SeekableStream
masterOff int64
readerMap sync.Map
readers orderedReaders
headCache *headCache
}

type orderedReaders struct {
mu sync.Mutex
m map[int64]io.Reader
keys []int64
}

func (o *orderedReaders) store(off int64, r io.Reader) {
o.mu.Lock()
defer o.mu.Unlock()
if _, ok := o.m[off]; ok {
o.m[off] = r
return
}
if o.m == nil {
o.m = make(map[int64]io.Reader)
}
i := sort.Search(len(o.keys), func(i int) bool { return o.keys[i] >= off })
o.keys = append(o.keys, 0)
copy(o.keys[i+1:], o.keys[i:])
o.keys[i] = off
o.m[off] = r
}

func (o *orderedReaders) takeExact(off int64) (io.Reader, bool) {
o.mu.Lock()
defer o.mu.Unlock()
r, ok := o.m[off]
if ok {
delete(o.m, off)
o.removeKey(off)
}
return r, ok
}

func (o *orderedReaders) takeBest(off int64) (io.Reader, int64, bool) {
o.mu.Lock()
defer o.mu.Unlock()
if r, ok := o.m[off]; ok {
delete(o.m, off)
o.removeKey(off)
return r, off, true
}
i := sort.Search(len(o.keys), func(i int) bool { return o.keys[i] >= off })
if i == 0 {
return nil, 0, false
}
k := o.keys[i-1]
if off-k > 4*utils.MB {
return nil, 0, false
}
r := o.m[k]
delete(o.m, k)
o.removeKey(k)
return r, k, true
}

func (o *orderedReaders) removeKey(k int64) {
i := sort.Search(len(o.keys), func(i int) bool { return o.keys[i] >= k })
copy(o.keys[i:], o.keys[i+1:])
o.keys = o.keys[:len(o.keys)-1]
}

type headCache struct {
reader io.Reader
bufs [][]byte
Expand Down Expand Up @@ -396,7 +459,7 @@ func (r *headCache) Close() error {

func (r *RangeReadReadAtSeeker) InitHeadCache() {
if r.masterOff == 0 {
value, _ := r.readerMap.LoadAndDelete(int64(0))
value, _ := r.readers.takeExact(0)
r.headCache = &headCache{reader: value.(io.Reader)}
r.ss.Closers.Add(r.headCache)
}
Expand All @@ -422,9 +485,9 @@ func NewReadAtSeeker(ss *SeekableStream, offset int64, forceRange ...bool) (mode
if err != nil {
return nil, err
}
r.readerMap.Store(int64(offset), reader)
r.readers.store(offset, reader)
} else {
r.readerMap.Store(int64(offset), ss)
r.readers.store(0, ss)
}
return r, nil
}
Expand All @@ -442,41 +505,15 @@ func NewMultiReaderAt(ss []*SeekableStream) (readerutil.SizeReaderAt, error) {
}

func (r *RangeReadReadAtSeeker) getReaderAtOffset(off int64) (io.Reader, error) {
for {
var cur int64 = -1
r.readerMap.Range(func(key, value any) bool {
k := key.(int64)
if off == k {
cur = k
return false
}
if off > k && off-k <= 4*utils.MB && k > cur {
cur = k
}
return true
})
if cur < 0 {
break
}
v, ok := r.readerMap.LoadAndDelete(int64(cur))
if !ok {
continue
}
rr := v.(io.Reader)
if off == int64(cur) {
// logrus.Debugf("getReaderAtOffset match_%d", off)
if rr, cur, ok := r.readers.takeBest(off); ok {
if cur == off {
return rr, nil
}
n, _ := utils.CopyWithBufferN(io.Discard, rr, off-cur)
cur += n
if cur == off {
// logrus.Debugf("getReaderAtOffset old_%d", off)
if cur+n == off {
return rr, nil
}
break
}

// logrus.Debugf("getReaderAtOffset new_%d", off)
reader, err := r.ss.RangeRead(http_range.Range{Start: off, Length: -1})
if err != nil {
return nil, err
Expand All @@ -501,7 +538,7 @@ func (r *RangeReadReadAtSeeker) ReadAt(p []byte, off int64) (n int, err error) {
off += int64(n)
switch err {
case nil:
r.readerMap.Store(int64(off), rr)
r.readers.store(off, rr)
case io.ErrUnexpectedEOF:
err = io.EOF
}
Expand Down
Loading