mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 17:50:12 +00:00
fix(rpc): make IO stream lifecycle race-safe
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user