fix(rpc): make IO stream lifecycle race-safe

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-07-20 04:26:19 +00:00
co-authored by naiba/CloudCode
parent 8b47ff141f
commit c756ef9385
23 changed files with 2114 additions and 690 deletions
+28 -7
View File
@@ -6,6 +6,7 @@ import (
"io"
"net"
"net/http"
"sync"
)
var _ io.ReadWriteCloser = (*RequestWrapper)(nil)
@@ -14,6 +15,11 @@ type RequestWrapper struct {
req *http.Request
reader *bytes.Buffer
writer net.Conn
closeOnce sync.Once
closeInit sync.Once
closeDone chan struct{}
closeErr error
}
func NewRequestWrapper(req *http.Request, writer http.ResponseWriter) (*RequestWrapper, error) {
@@ -27,12 +33,17 @@ func NewRequestWrapper(req *http.Request, writer http.ResponseWriter) (*RequestW
}
buf := bytes.NewBuffer(nil)
if err = req.Write(buf); err != nil {
return nil, err
var bodyErr error
if req.Body != nil {
bodyErr = req.Body.Close()
}
return nil, errors.Join(err, bodyErr, conn.Close())
}
return &RequestWrapper{
req: req,
reader: buf,
writer: conn,
req: req,
reader: buf,
writer: conn,
closeDone: make(chan struct{}),
}, nil
}
@@ -53,7 +64,17 @@ func (rw *RequestWrapper) Write(p []byte) (int, error) {
}
func (rw *RequestWrapper) Close() error {
rw.req.Body.Close()
rw.writer.Close()
return nil
rw.closeInit.Do(func() {
rw.closeDone = make(chan struct{})
})
rw.closeOnce.Do(func() {
var bodyErr error
if rw.req.Body != nil {
bodyErr = rw.req.Body.Close()
}
rw.closeErr = errors.Join(bodyErr, rw.writer.Close())
close(rw.closeDone)
})
<-rw.closeDone
return rw.closeErr
}
+232
View File
@@ -0,0 +1,232 @@
package utils
import (
"bufio"
"bytes"
"errors"
"io"
"net"
"net/http"
"net/url"
"strings"
"sync"
"sync/atomic"
"testing"
)
var (
errRequestWrite = errors.New("request write failed")
errBodyClose = errors.New("body close failed")
errConnClose = errors.New("connection close failed")
)
func TestNewRequestWrapper_closesHijackedConnectionWhenRequestWriteFails(t *testing.T) {
// Given
conn := newRequestWrapperTestConn(errConnClose)
req := &http.Request{
Method: "POST",
URL: &url.URL{Scheme: "http", Host: "example.test", Path: "/nat"},
Body: &requestWrapperTestBody{readErr: errRequestWrite, closeErr: errBodyClose},
ContentLength: 1,
}
writer := &requestWrapperTestResponseWriter{conn: conn}
// When
_, err := NewRequestWrapper(req, writer)
// Then
if err == nil || !strings.Contains(err.Error(), errRequestWrite.Error()) {
t.Fatalf("expected request write error, got %v", err)
}
if !errors.Is(err, errBodyClose) {
t.Fatalf("expected body close error, got %v", err)
}
if !errors.Is(err, errConnClose) {
t.Fatalf("expected connection close error, got %v", err)
}
if got := conn.closeCount.Load(); got != 1 {
t.Fatalf("expected one connection close, got %d", got)
}
}
func TestRequestWrapper_Close_joinsBodyAndConnectionErrors(t *testing.T) {
// Given
body := &requestWrapperTestBody{closeErr: errBodyClose}
conn := newRequestWrapperTestConn(errConnClose)
rw := &RequestWrapper{
req: &http.Request{Body: body},
reader: bytes.NewBuffer(nil),
writer: conn,
closeDone: make(chan struct{}),
}
// When
err := rw.Close()
// Then
if !errors.Is(err, errBodyClose) {
t.Fatalf("expected body close error, got %v", err)
}
if !errors.Is(err, errConnClose) {
t.Fatalf("expected connection close error, got %v", err)
}
}
func TestRequestWrapper_Close_repeatedCallersReceiveRetainedErrorAndCloseOnce(t *testing.T) {
// Given
body := &requestWrapperTestBody{closeErr: errBodyClose}
conn := newRequestWrapperTestConn(errConnClose)
rw := &RequestWrapper{
req: &http.Request{Body: body},
reader: bytes.NewBuffer(nil),
writer: conn,
closeDone: make(chan struct{}),
}
// When
firstErr := rw.Close()
secondErr := rw.Close()
// Then
if firstErr != secondErr {
t.Fatalf("expected identical retained error, got distinct values %p and %p", firstErr, secondErr)
}
if got := body.closeCount.Load(); got != 1 {
t.Fatalf("expected one body close, got %d", got)
}
if got := conn.closeCount.Load(); got != 1 {
t.Fatalf("expected one connection close, got %d", got)
}
}
func TestRequestWrapper_Close_concurrentCallersWaitForCopyToUnblock(t *testing.T) {
// Given
body := &requestWrapperTestBody{closeErr: errBodyClose}
conn := newRequestWrapperTestConn(errConnClose)
rw := &RequestWrapper{
req: &http.Request{Body: body},
reader: bytes.NewBuffer(nil),
writer: conn,
closeDone: make(chan struct{}),
}
readDone := make(chan error, 1)
go func() {
_, err := rw.Read(make([]byte, 1))
readDone <- err
}()
<-conn.readStarted
// When
const callerCount = 8
results := make(chan error, callerCount)
var callers sync.WaitGroup
callers.Add(callerCount)
for range callerCount {
go func() {
defer callers.Done()
results <- rw.Close()
}()
}
callers.Wait()
close(results)
// Then
var retainedErr error
for err := range results {
if !errors.Is(err, errConnClose) {
t.Fatalf("expected retained connection close error, got %v", err)
}
if !errors.Is(err, errBodyClose) {
t.Fatalf("expected retained body close error, got %v", err)
}
if retainedErr == nil {
retainedErr = err
continue
}
if err != retainedErr {
t.Fatalf("expected identical retained error, got distinct values %p and %p", retainedErr, err)
}
}
<-readDone
if got := body.closeCount.Load(); got != 1 {
t.Fatalf("expected one body close, got %d", got)
}
if got := conn.closeCount.Load(); got != 1 {
t.Fatalf("expected one connection close, got %d", got)
}
}
type requestWrapperTestBody struct {
readErr error
closeErr error
closeCount atomic.Int32
}
func (b *requestWrapperTestBody) Read([]byte) (int, error) {
if b.readErr != nil {
return 0, b.readErr
}
return 0, io.EOF
}
func (b *requestWrapperTestBody) Close() error {
b.closeCount.Add(1)
return b.closeErr
}
type requestWrapperTestConn struct {
net.Conn
closeErr error
closeCount atomic.Int32
closeOnce sync.Once
readStarted chan struct{}
readDone chan struct{}
}
func newRequestWrapperTestConn(closeErr error) *requestWrapperTestConn {
return &requestWrapperTestConn{
closeErr: closeErr,
readStarted: make(chan struct{}),
readDone: make(chan struct{}),
}
}
func (c *requestWrapperTestConn) Read([]byte) (int, error) {
select {
case <-c.readStarted:
default:
close(c.readStarted)
}
<-c.readDone
return 0, io.ErrClosedPipe
}
func (c *requestWrapperTestConn) Write(p []byte) (int, error) {
return len(p), nil
}
func (c *requestWrapperTestConn) Close() error {
c.closeOnce.Do(func() {
c.closeCount.Add(1)
close(c.readDone)
})
return c.closeErr
}
type requestWrapperTestResponseWriter struct {
conn net.Conn
}
func (w *requestWrapperTestResponseWriter) Header() http.Header {
return make(http.Header)
}
func (w *requestWrapperTestResponseWriter) Write([]byte) (int, error) {
return 0, nil
}
func (w *requestWrapperTestResponseWriter) WriteHeader(int) {}
func (w *requestWrapperTestResponseWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
return w.conn, bufio.NewReadWriter(bufio.NewReader(bytes.NewReader(nil)), bufio.NewWriter(io.Discard)), nil
}