From 05e7232d8fb745c24253da44c727efbf8daffcc7 Mon Sep 17 00:00:00 2001 From: sususu Date: Sun, 2 Aug 2026 18:00:57 +0800 Subject: [PATCH] fix(httpwire): retain request state after partial writes --- internal/httpwire/ordered_conn.go | 160 +++++++++++++++++++++---- internal/httpwire/ordered_conn_test.go | 110 ++++++++++++++++- 2 files changed, 240 insertions(+), 30 deletions(-) diff --git a/internal/httpwire/ordered_conn.go b/internal/httpwire/ordered_conn.go index 77386f471..4b78aefa1 100644 --- a/internal/httpwire/ordered_conn.go +++ b/internal/httpwire/ordered_conn.go @@ -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 } diff --git a/internal/httpwire/ordered_conn_test.go b/internal/httpwire/ordered_conn_test.go index 02fbb4f97..eb9fbe846 100644 --- a/internal/httpwire/ordered_conn_test.go +++ b/internal/httpwire/ordered_conn_test.go @@ -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) + } +}