Skip to content

Commit b68d71c

Browse files
committed
memsys: guard SGL vs lingering HTTP request-body reads
* `net/http` may keep reading (and close) a request body after `client.Do` returns in cases such as proxy redirect or retry * freeing the part-upload SGL at that point is use-after-free: panic in the read path or , worse, slient corruption via recycled slab buffers * this patch add `GuardReader`: reads hold the guard shared, `Free` takes it exclusively and waits out in-flight reads Signed-off-by: Tony Chen <a122774007@gmail.com>
1 parent 7ad0be7 commit b68d71c

3 files changed

Lines changed: 219 additions & 8 deletions

File tree

ais/tgtmpt.go

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -358,16 +358,17 @@ func (ups *ups) _put(args *partArgs) (etag string, ecode int, err error) {
358358
sgl := t.gmm.NewSGL(rsize)
359359
mw.Append(sgl)
360360
expectedSize, err = io.Copy(mw, reader)
361+
// NOTE: memsys reader wrapper required:
362+
// - Azure: seekable (io.ReadSeekCloser) for SDK's seek-based retry
363+
// - remote AIS: retriable (cos.ReadOpenCloser) for DoWithRetry (reader.Open() on retry)
364+
// - guarded: the transport may still be reading the body when PutMptPart returns
365+
rdr := memsys.NewGuardReader(sgl)
361366
if err == nil {
362-
// NOTE: memsys.Reader wrapper required:
363-
// - Azure: seekable (io.ReadSeekCloser) for SDK's seek-based retry
364-
// - Remote AIS: retriable (cos.ReadOpenCloser) for DoWithRetry (reader.Open() on retry)
365-
rdr := memsys.NewReader(sgl)
366367
remoteStart := mono.NanoTime()
367368
etag, ecode, err = backend.PutMptPart(lom, rdr, args.req, uploadID, expectedSize, int32(args.partNum))
368369
remotePutLatency = mono.SinceNano(remoteStart)
369370
}
370-
sgl.Free()
371+
rdr.Free() // not sgl.Free
371372
default:
372373
// high memory pressure
373374
time.Sleep(cos.PollSleepLong) // throttle
@@ -383,13 +384,13 @@ func (ups *ups) _put(args *partArgs) (etag string, ecode int, err error) {
383384
sgl := t.gmm.NewSGL(rsize)
384385
mw.Append(sgl)
385386
expectedSize, err = io.Copy(mw, reader)
387+
rdr := memsys.NewGuardReader(sgl)
386388
if err == nil {
387-
rdr := memsys.NewReader(sgl)
388389
remoteStart := mono.NanoTime()
389390
etag, ecode, err = backend.PutMptPart(lom, rdr, args.req, uploadID, expectedSize, int32(args.partNum))
390391
remotePutLatency = mono.SinceNano(remoteStart)
391392
}
392-
sgl.Free()
393+
rdr.Free() // not sgl.Free (ditto)
393394
}
394395

395396
// Release streaming checksum lock

memsys/guardreader_test.go

Lines changed: 142 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,142 @@
1+
// Package memsys provides memory management and Slab allocation
2+
// with io.Reader and io.Writer interfaces on top of a scatter-gather lists
3+
// (of reusable buffers)
4+
/*
5+
* Copyright (c) 2026, NVIDIA CORPORATION. All rights reserved.
6+
*/
7+
package memsys_test
8+
9+
import (
10+
"fmt"
11+
"io"
12+
"sync"
13+
"testing"
14+
15+
"github.com/NVIDIA/aistore/memsys"
16+
"github.com/NVIDIA/aistore/tools/tassert"
17+
)
18+
19+
// GuardReader reads are guarded vs Free; reads-after-Free return EOF; see memsys.GuardReader
20+
func TestGuardReaderFree(t *testing.T) {
21+
mem := &memsys.MMSA{Name: "guardr.test", MinPctFree: 50}
22+
mem.Init(0)
23+
defer mem.Terminate(false)
24+
25+
sgl := mem.NewSGL(128)
26+
sgl.Write(make([]byte, 3000))
27+
rdr := memsys.NewGuardReader(sgl)
28+
29+
view, err := rdr.Open()
30+
tassert.CheckFatal(t, err)
31+
32+
n, err := view.Read(make([]byte, 1000))
33+
tassert.Fatalf(t, n == 1000 && err == nil, "read before Free: (%d, %v)", n, err)
34+
tassert.Fatalf(t, !sgl.IsNil(), "SGL freed before Free")
35+
36+
rdr.Free()
37+
rdr.Free() // idempotent
38+
tassert.Fatalf(t, sgl.IsNil(), "SGL must be freed upon GuardReader.Free")
39+
40+
// lingering transport reads after Free: clean EOF, no panic
41+
n, err = view.Read(make([]byte, 1000))
42+
tassert.Fatalf(t, n == 0 && err == io.EOF, "read after Free: expected (0, EOF), got (%d, %v)", n, err)
43+
}
44+
45+
// concurrent reads racing Free: no panic, no unexpected error, no SGL leak
46+
func TestGuardReaderStress(t *testing.T) {
47+
mem := &memsys.MMSA{Name: "guardr.stress", MinPctFree: 50}
48+
mem.Init(0)
49+
defer mem.Terminate(false)
50+
51+
const (
52+
numReaders = 4
53+
rounds = 100
54+
)
55+
for range rounds {
56+
sgl := mem.NewSGL(128)
57+
sgl.Write(make([]byte, 8192))
58+
rdr := memsys.NewGuardReader(sgl)
59+
60+
var (
61+
wg sync.WaitGroup
62+
errCh = make(chan error, numReaders)
63+
started = make(chan struct{}, numReaders)
64+
)
65+
for range numReaders {
66+
view, err := rdr.Open()
67+
tassert.CheckFatal(t, err)
68+
wg.Add(1)
69+
go func() {
70+
defer wg.Done()
71+
started <- struct{}{}
72+
buf := make([]byte, 512)
73+
for {
74+
_, err := view.Read(buf)
75+
if err == io.EOF {
76+
return
77+
}
78+
if err != nil {
79+
errCh <- fmt.Errorf("unexpected read error: %v", err)
80+
return
81+
}
82+
}
83+
}()
84+
}
85+
for range numReaders {
86+
<-started
87+
}
88+
rdr.Free() // concurrently with the readers
89+
rdr.Free() // idempotent, ditto
90+
wg.Wait()
91+
close(errCh)
92+
for err := range errCh {
93+
tassert.CheckFatal(t, err)
94+
}
95+
tassert.Fatalf(t, sgl.IsNil(), "SGL leaked with concurrent readers racing Free")
96+
}
97+
}
98+
99+
// deterministic overlap: reader completes one read, signals, and keeps reading
100+
// while main frees mid-stream; the reader must then observe clean reads or EOF
101+
func TestGuardReaderFreeMidStream(t *testing.T) {
102+
mem := &memsys.MMSA{Name: "guardr.midstream", MinPctFree: 50}
103+
mem.Init(0)
104+
defer mem.Terminate(false)
105+
106+
sgl := mem.NewSGL(128)
107+
sgl.Write(make([]byte, 8192))
108+
rdr := memsys.NewGuardReader(sgl)
109+
110+
view, err := rdr.Open()
111+
tassert.CheckFatal(t, err)
112+
113+
var (
114+
firstRead = make(chan struct{})
115+
done = make(chan error, 1)
116+
)
117+
go func() {
118+
buf := make([]byte, 512)
119+
n, err := view.Read(buf)
120+
if n != len(buf) || err != nil {
121+
done <- fmt.Errorf("first read: (%d, %v)", n, err)
122+
return
123+
}
124+
close(firstRead) // signal main to Free
125+
for {
126+
_, err := view.Read(buf) // keep reading, racing Free
127+
if err == io.EOF {
128+
done <- nil
129+
return
130+
}
131+
if err != nil {
132+
done <- fmt.Errorf("unexpected read error: %v", err)
133+
return
134+
}
135+
}
136+
}()
137+
138+
<-firstRead
139+
rdr.Free()
140+
tassert.CheckFatal(t, <-done)
141+
tassert.Fatalf(t, sgl.IsNil(), "SGL must be freed after mid-stream Free")
142+
}

memsys/iosgl.go

Lines changed: 69 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,7 @@
11
// Package memsys provides memory management and slab/SGL allocation with io.Reader and io.Writer interfaces
22
// on top of scatter-gather lists of reusable buffers.
33
/*
4-
* Copyright (c) 2018-2025, NVIDIA CORPORATION. All rights reserved.
4+
* Copyright (c) 2018-2026, NVIDIA CORPORATION. All rights reserved.
55
*/
66
package memsys
77

@@ -10,6 +10,7 @@ import (
1010
"errors"
1111
"fmt"
1212
"io"
13+
"io/fs"
1314
"sync"
1415

1516
"github.com/NVIDIA/aistore/cmn"
@@ -31,6 +32,9 @@ var (
3132
_ cos.ReadOpenCloser = (*SGL)(nil)
3233
_ cos.ReadOpenCloser = (*Reader)(nil)
3334
_ io.ReadSeekCloser = (*Reader)(nil)
35+
36+
_ cos.ReadOpenCloser = (*GuardReader)(nil)
37+
_ io.ReadSeekCloser = (*GuardReader)(nil)
3438
)
3539

3640
type (
@@ -419,6 +423,70 @@ func (r *Reader) Read(b []byte) (n int, err error) {
419423
return n, err
420424
}
421425

426+
/////////////////
427+
// GuardReader //
428+
/////////////////
429+
430+
// GuardReader is a Reader to use when the underlying SGL may get freed while
431+
// readers are still around.
432+
//
433+
// Reads hold the guard shared; Free takes it exclusively - waiting out any
434+
// read in flight - marks, and frees the SGL; subsequent reads return io.EOF.
435+
type (
436+
guard struct {
437+
sync.RWMutex
438+
freed bool // only ever accessed under the lock
439+
}
440+
GuardReader struct {
441+
Reader
442+
g *guard // shared by all Open-ed views (each view has its own read offset)
443+
}
444+
)
445+
446+
func NewGuardReader(z *SGL) *GuardReader { return &GuardReader{Reader{z, 0}, &guard{}} }
447+
448+
func (r *GuardReader) Open() (cos.ReadOpenCloser, error) {
449+
return &GuardReader{Reader{r.z, 0}, r.g}, nil
450+
}
451+
452+
func (r *GuardReader) Read(b []byte) (n int, err error) {
453+
r.g.RLock()
454+
if r.g.freed {
455+
r.g.RUnlock()
456+
return 0, io.EOF
457+
}
458+
n, err = r.Reader.Read(b)
459+
r.g.RUnlock()
460+
return n, err
461+
}
462+
463+
// guarded as well: SeekEnd reads SGL state
464+
func (r *GuardReader) Seek(from int64, whence int) (offset int64, err error) {
465+
r.g.RLock()
466+
if r.g.freed {
467+
r.g.RUnlock()
468+
return 0, fs.ErrClosed
469+
}
470+
offset, err = r.Reader.Seek(from, whence)
471+
r.g.RUnlock()
472+
return offset, err
473+
}
474+
475+
// Close is a no-op: closing a view does not free the SGL - Free does
476+
func (*GuardReader) Close() error { return nil }
477+
478+
// Free the underlying SGL; idempotent
479+
// (frees after Unlock - never holds two locks)
480+
func (r *GuardReader) Free() {
481+
r.g.Lock()
482+
freed := r.g.freed
483+
r.g.freed = true
484+
r.g.Unlock()
485+
if !freed {
486+
r.z.Free()
487+
}
488+
}
489+
422490
func (r *Reader) Seek(from int64, whence int) (offset int64, err error) {
423491
switch whence {
424492
case io.SeekStart:

0 commit comments

Comments
 (0)