mirror of
https://github.com/router-for-me/CLIProxyAPI.git
synced 2026-09-03 06:35:00 +08:00
fix(httpwire): retain request state after partial writes
This commit is contained in:
@@ -34,38 +34,56 @@ type orderedRequestConn struct {
|
||||
mu sync.Mutex
|
||||
header []byte
|
||||
bodyRemaining int64
|
||||
passthrough bool
|
||||
chunked *chunkedRequestTracker
|
||||
}
|
||||
|
||||
func (c *orderedRequestConn) Write(p []byte) (int, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
if c.passthrough {
|
||||
return c.Conn.Write(p)
|
||||
}
|
||||
|
||||
originalLength := len(p)
|
||||
consumed := 0
|
||||
remaining := p
|
||||
for len(remaining) > 0 {
|
||||
if c.bodyRemaining > 0 {
|
||||
bodyBytes := int64(len(remaining))
|
||||
if bodyBytes > c.bodyRemaining {
|
||||
bodyBytes = c.bodyRemaining
|
||||
}
|
||||
if errWrite := writeAll(c.Conn, remaining[:bodyBytes]); errWrite != nil {
|
||||
return 0, errWrite
|
||||
bodyBytes := min(int64(len(remaining)), c.bodyRemaining)
|
||||
written, errWrite := writeAll(c.Conn, remaining[:bodyBytes])
|
||||
consumed += written
|
||||
c.bodyRemaining -= int64(written)
|
||||
if errWrite != nil {
|
||||
return consumed, errWrite
|
||||
}
|
||||
remaining = remaining[bodyBytes:]
|
||||
c.bodyRemaining -= bodyBytes
|
||||
continue
|
||||
}
|
||||
if c.chunked != nil {
|
||||
preview := c.chunked.clone()
|
||||
chunkBytes, _, errChunk := preview.consume(remaining)
|
||||
if errChunk != nil {
|
||||
return consumed, errChunk
|
||||
}
|
||||
written, errWrite := writeAll(c.Conn, remaining[:chunkBytes])
|
||||
consumed += written
|
||||
_, completed, errConsume := c.chunked.consume(remaining[:written])
|
||||
if errConsume != nil {
|
||||
return consumed, errConsume
|
||||
}
|
||||
if completed {
|
||||
c.chunked = nil
|
||||
}
|
||||
if errWrite != nil {
|
||||
return consumed, errWrite
|
||||
}
|
||||
remaining = remaining[chunkBytes:]
|
||||
continue
|
||||
}
|
||||
|
||||
previousHeaderLength := len(c.header)
|
||||
c.header = append(c.header, remaining...)
|
||||
headerEnd := bytes.Index(c.header, []byte("\r\n\r\n"))
|
||||
if headerEnd < 0 {
|
||||
if len(c.header) > maxBufferedRequestHeader {
|
||||
return 0, fmt.Errorf("httpwire: request header exceeds %d bytes", maxBufferedRequestHeader)
|
||||
return consumed, fmt.Errorf("httpwire: request header exceeds %d bytes", maxBufferedRequestHeader)
|
||||
}
|
||||
return originalLength, nil
|
||||
}
|
||||
@@ -74,20 +92,22 @@ func (c *orderedRequestConn) Write(p []byte) (int, error) {
|
||||
header := c.header[:headerEnd]
|
||||
body := c.header[headerEnd:]
|
||||
c.header = nil
|
||||
currentHeaderBytes := min(len(remaining), max(0, headerEnd-previousHeaderLength))
|
||||
|
||||
ordered, contentLength, chunked := orderRequestHeader(header, c.order)
|
||||
if errWrite := writeAll(c.Conn, ordered); errWrite != nil {
|
||||
return 0, errWrite
|
||||
if _, errWrite := writeAll(c.Conn, ordered); errWrite != nil {
|
||||
// All caller bytes were accepted into the wrapper before the transformed
|
||||
// header write failed. Return the full input count with the terminal
|
||||
// connection error so callers do not replay an ambiguous partial header.
|
||||
return originalLength, errWrite
|
||||
}
|
||||
consumed += currentHeaderBytes
|
||||
remaining = body
|
||||
if chunked {
|
||||
if errWrite := writeAll(c.Conn, body); errWrite != nil {
|
||||
return 0, errWrite
|
||||
}
|
||||
c.passthrough = true
|
||||
return originalLength, nil
|
||||
c.chunked = newChunkedRequestTracker()
|
||||
continue
|
||||
}
|
||||
c.bodyRemaining = contentLength
|
||||
remaining = body
|
||||
}
|
||||
return originalLength, nil
|
||||
}
|
||||
@@ -171,16 +191,106 @@ func requestUsesChunkedEncoding(lines [][]byte) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
func writeAll(writer io.Writer, data []byte) error {
|
||||
type chunkedRequestTracker struct {
|
||||
state uint8
|
||||
line []byte
|
||||
dataRemaining int64
|
||||
crlfPosition int
|
||||
trailers []byte
|
||||
}
|
||||
|
||||
const (
|
||||
chunkedReadingSize uint8 = iota
|
||||
chunkedReadingData
|
||||
chunkedReadingDataCRLF
|
||||
chunkedReadingTrailers
|
||||
)
|
||||
|
||||
func newChunkedRequestTracker() *chunkedRequestTracker {
|
||||
return &chunkedRequestTracker{state: chunkedReadingSize}
|
||||
}
|
||||
|
||||
func (tracker *chunkedRequestTracker) clone() *chunkedRequestTracker {
|
||||
cloned := *tracker
|
||||
cloned.line = append([]byte(nil), tracker.line...)
|
||||
cloned.trailers = append([]byte(nil), tracker.trailers...)
|
||||
return &cloned
|
||||
}
|
||||
|
||||
func (tracker *chunkedRequestTracker) consume(data []byte) (consumed int, completed bool, err error) {
|
||||
for consumed < len(data) {
|
||||
switch tracker.state {
|
||||
case chunkedReadingSize:
|
||||
tracker.line = append(tracker.line, data[consumed])
|
||||
consumed++
|
||||
if len(tracker.line) > maxBufferedRequestHeader {
|
||||
return consumed, false, fmt.Errorf("httpwire: chunk size line exceeds %d bytes", maxBufferedRequestHeader)
|
||||
}
|
||||
if len(tracker.line) < 2 || !bytes.Equal(tracker.line[len(tracker.line)-2:], []byte("\r\n")) {
|
||||
continue
|
||||
}
|
||||
sizeText := strings.TrimSpace(string(tracker.line[:len(tracker.line)-2]))
|
||||
if extension := strings.IndexByte(sizeText, ';'); extension >= 0 {
|
||||
sizeText = strings.TrimSpace(sizeText[:extension])
|
||||
}
|
||||
size, errParse := strconv.ParseInt(sizeText, 16, 64)
|
||||
if errParse != nil || size < 0 {
|
||||
return consumed, false, fmt.Errorf("httpwire: invalid chunk size %q", sizeText)
|
||||
}
|
||||
tracker.line = tracker.line[:0]
|
||||
if size == 0 {
|
||||
tracker.state = chunkedReadingTrailers
|
||||
continue
|
||||
}
|
||||
tracker.dataRemaining = size
|
||||
tracker.state = chunkedReadingData
|
||||
case chunkedReadingData:
|
||||
chunkBytes := min(int64(len(data)-consumed), tracker.dataRemaining)
|
||||
consumed += int(chunkBytes)
|
||||
tracker.dataRemaining -= chunkBytes
|
||||
if tracker.dataRemaining == 0 {
|
||||
tracker.crlfPosition = 0
|
||||
tracker.state = chunkedReadingDataCRLF
|
||||
}
|
||||
case chunkedReadingDataCRLF:
|
||||
want := []byte("\r\n")
|
||||
if data[consumed] != want[tracker.crlfPosition] {
|
||||
return consumed, false, fmt.Errorf("httpwire: chunk data is missing CRLF terminator")
|
||||
}
|
||||
consumed++
|
||||
tracker.crlfPosition++
|
||||
if tracker.crlfPosition == len(want) {
|
||||
tracker.state = chunkedReadingSize
|
||||
}
|
||||
case chunkedReadingTrailers:
|
||||
tracker.trailers = append(tracker.trailers, data[consumed])
|
||||
consumed++
|
||||
if len(tracker.trailers) > maxBufferedRequestHeader {
|
||||
return consumed, false, fmt.Errorf("httpwire: chunk trailers exceed %d bytes", maxBufferedRequestHeader)
|
||||
}
|
||||
if bytes.Equal(tracker.trailers, []byte("\r\n")) ||
|
||||
(len(tracker.trailers) >= 4 && bytes.Equal(tracker.trailers[len(tracker.trailers)-4:], []byte("\r\n\r\n"))) {
|
||||
return consumed, true, nil
|
||||
}
|
||||
default:
|
||||
return consumed, false, fmt.Errorf("httpwire: invalid chunk parser state %d", tracker.state)
|
||||
}
|
||||
}
|
||||
return consumed, false, nil
|
||||
}
|
||||
|
||||
func writeAll(writer io.Writer, data []byte) (int, error) {
|
||||
total := 0
|
||||
for len(data) > 0 {
|
||||
written, errWrite := writer.Write(data)
|
||||
total += written
|
||||
if errWrite != nil {
|
||||
return errWrite
|
||||
return total, errWrite
|
||||
}
|
||||
if written <= 0 {
|
||||
return io.ErrShortWrite
|
||||
return total, io.ErrShortWrite
|
||||
}
|
||||
data = data[written:]
|
||||
}
|
||||
return nil
|
||||
return total, nil
|
||||
}
|
||||
|
||||
@@ -73,7 +73,7 @@ func TestOrderedRequestConnReordersKeepAliveRequestsWithoutChangingBodies(t *tes
|
||||
}
|
||||
}
|
||||
|
||||
func TestOrderedRequestConnPreservesChunkedBody(t *testing.T) {
|
||||
func TestOrderedRequestConnPreservesChunkedBodyAndReordersNextRequest(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
client, server := net.Pipe()
|
||||
@@ -82,8 +82,10 @@ func TestOrderedRequestConnPreservesChunkedBody(t *testing.T) {
|
||||
_ = server.Close()
|
||||
})
|
||||
conn := NewOrderedRequestConn(client, func(_, _ string) []string { return []string{"Host", "Transfer-Encoding"} })
|
||||
input := []byte("POST /upload HTTP/1.1\r\nTransfer-Encoding: chunked\r\nHost: example.com\r\n\r\n4\r\ntest\r\n0\r\n\r\n")
|
||||
want := []byte("POST /upload HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\n\r\n4\r\ntest\r\n0\r\n\r\n")
|
||||
first := "POST /upload HTTP/1.1\r\nTransfer-Encoding: chunked\r\nHost: example.com\r\n\r\n4\r\ntest\r\n0\r\nX-Trailer: done\r\n\r\n"
|
||||
second := "GET /next HTTP/1.1\r\nTransfer-Encoding: identity\r\nHost: example.com\r\n\r\n"
|
||||
input := []byte(first + second)
|
||||
want := []byte("POST /upload HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\n\r\n4\r\ntest\r\n0\r\nX-Trailer: done\r\n\r\nGET /next HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: identity\r\n\r\n")
|
||||
|
||||
readDone := make(chan []byte, 1)
|
||||
go func() {
|
||||
@@ -91,10 +93,108 @@ func TestOrderedRequestConnPreservesChunkedBody(t *testing.T) {
|
||||
_, _ = io.ReadFull(server, got)
|
||||
readDone <- got
|
||||
}()
|
||||
if _, errWrite := conn.Write(input); errWrite != nil {
|
||||
t.Fatal(errWrite)
|
||||
for index := range input {
|
||||
part := input[index : index+1]
|
||||
written, errWrite := conn.Write(part)
|
||||
if errWrite != nil {
|
||||
t.Fatal(errWrite)
|
||||
}
|
||||
if written != len(part) {
|
||||
t.Fatalf("write length = %d, want %d", written, len(part))
|
||||
}
|
||||
}
|
||||
if got := <-readDone; !bytes.Equal(got, want) {
|
||||
t.Fatalf("chunked wire bytes differ\n got: %q\nwant: %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
type partialErrorConn struct {
|
||||
bytes.Buffer
|
||||
failLimit int
|
||||
failErr error
|
||||
}
|
||||
|
||||
func (conn *partialErrorConn) Write(data []byte) (int, error) {
|
||||
if conn.failErr == nil {
|
||||
return conn.Buffer.Write(data)
|
||||
}
|
||||
written := min(conn.failLimit, len(data))
|
||||
_, _ = conn.Buffer.Write(data[:written])
|
||||
return written, conn.failErr
|
||||
}
|
||||
|
||||
func (*partialErrorConn) Read([]byte) (int, error) { return 0, io.EOF }
|
||||
func (*partialErrorConn) Close() error { return nil }
|
||||
func (*partialErrorConn) LocalAddr() net.Addr { return nil }
|
||||
func (*partialErrorConn) RemoteAddr() net.Addr { return nil }
|
||||
func (*partialErrorConn) SetDeadline(time.Time) error { return nil }
|
||||
func (*partialErrorConn) SetReadDeadline(time.Time) error { return nil }
|
||||
func (*partialErrorConn) SetWriteDeadline(time.Time) error { return nil }
|
||||
|
||||
func TestOrderedRequestConnReportsPartialBodyWrite(t *testing.T) {
|
||||
underlying := &partialErrorConn{}
|
||||
conn := NewOrderedRequestConn(underlying, func(_, _ string) []string { return []string{"Host", "Content-Length"} })
|
||||
header := []byte("POST /upload HTTP/1.1\r\nContent-Length: 5\r\nHost: example.com\r\n\r\n")
|
||||
if written, errWrite := conn.Write(header); errWrite != nil || written != len(header) {
|
||||
t.Fatalf("header write = %d, %v", written, errWrite)
|
||||
}
|
||||
|
||||
underlying.failLimit = 2
|
||||
injectedErr := errors.New("injected partial write")
|
||||
underlying.failErr = injectedErr
|
||||
written, errWrite := conn.Write([]byte("hello"))
|
||||
if !errors.Is(errWrite, injectedErr) {
|
||||
t.Fatalf("body write error = %v, want injected error", errWrite)
|
||||
}
|
||||
if written != 2 {
|
||||
t.Fatalf("body write length = %d, want underlying partial count 2", written)
|
||||
}
|
||||
if remaining := conn.(*orderedRequestConn).bodyRemaining; remaining != 3 {
|
||||
t.Fatalf("bodyRemaining = %d, want 3 after confirmed partial write", remaining)
|
||||
}
|
||||
|
||||
underlying.failErr = nil
|
||||
if written, errWrite = conn.Write([]byte("llo")); errWrite != nil || written != 3 {
|
||||
t.Fatalf("retried body write = %d, %v", written, errWrite)
|
||||
}
|
||||
second := []byte("GET /next HTTP/1.1\r\nContent-Length: 0\r\nHost: example.com\r\n\r\n")
|
||||
if written, errWrite = conn.Write(second); errWrite != nil || written != len(second) {
|
||||
t.Fatalf("next request write = %d, %v", written, errWrite)
|
||||
}
|
||||
want := "POST /upload HTTP/1.1\r\nHost: example.com\r\nContent-Length: 5\r\n\r\nhelloGET /next HTTP/1.1\r\nHost: example.com\r\nContent-Length: 0\r\n\r\n"
|
||||
if got := underlying.String(); got != want {
|
||||
t.Fatalf("wire bytes differ after retry\n got: %q\nwant: %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOrderedRequestConnTracksOnlyWrittenChunkBytesAfterPartialError(t *testing.T) {
|
||||
underlying := &partialErrorConn{}
|
||||
conn := NewOrderedRequestConn(underlying, func(_, _ string) []string { return []string{"Host", "Transfer-Encoding"} })
|
||||
header := []byte("POST /upload HTTP/1.1\r\nTransfer-Encoding: chunked\r\nHost: example.com\r\n\r\n")
|
||||
if written, errWrite := conn.Write(header); errWrite != nil || written != len(header) {
|
||||
t.Fatalf("header write = %d, %v", written, errWrite)
|
||||
}
|
||||
|
||||
chunkedBody := []byte("4\r\ntest\r\n0\r\nX-Trailer: done\r\n\r\n")
|
||||
underlying.failLimit = 6
|
||||
injectedErr := errors.New("injected chunk partial write")
|
||||
underlying.failErr = injectedErr
|
||||
written, errWrite := conn.Write(chunkedBody)
|
||||
if !errors.Is(errWrite, injectedErr) || written != 6 {
|
||||
t.Fatalf("chunk write = %d, %v; want 6 and injected error", written, errWrite)
|
||||
}
|
||||
|
||||
underlying.failErr = nil
|
||||
if retried, errRetry := conn.Write(chunkedBody[written:]); errRetry != nil || retried != len(chunkedBody)-written {
|
||||
t.Fatalf("retried chunk write = %d, %v", retried, errRetry)
|
||||
}
|
||||
second := []byte("GET /next HTTP/1.1\r\nTransfer-Encoding: identity\r\nHost: example.com\r\n\r\n")
|
||||
if written, errWrite = conn.Write(second); errWrite != nil || written != len(second) {
|
||||
t.Fatalf("next request write = %d, %v", written, errWrite)
|
||||
}
|
||||
want := "POST /upload HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: chunked\r\n\r\n" + string(chunkedBody) +
|
||||
"GET /next HTTP/1.1\r\nHost: example.com\r\nTransfer-Encoding: identity\r\n\r\n"
|
||||
if got := underlying.String(); got != want {
|
||||
t.Fatalf("wire bytes differ after chunk retry\n got: %q\nwant: %q", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user