From 27971f2a200a83d8d3d265be2be8686834af7786 Mon Sep 17 00:00:00 2001 From: niksis02 Date: Sat, 2 May 2026 01:07:07 +0400 Subject: [PATCH] fix: remove unsigned chunk reader caching Closes #1273 Rewrite UnsignedChunkReader to stream the payload bytes directly into the caller buffer instead of allocating and stashing full chunks. With this implementation, no stash is held by the reader and the chunk reader doesn't allocate any memory. Make debug logging more descriptive, which records reader state on all error paths and logs read progress whenever a Read call fills the caller buffer. Some unit tests were added to cover the main moving parts of the reader flow. --- s3api/utils/unsigned-chunk-reader.go | 356 +++++++++++++++------- s3api/utils/unsigned-chunk-reader_test.go | 177 +++++++++++ 2 files changed, 428 insertions(+), 105 deletions(-) create mode 100644 s3api/utils/unsigned-chunk-reader_test.go diff --git a/s3api/utils/unsigned-chunk-reader.go b/s3api/utils/unsigned-chunk-reader.go index be7b6d3b..e6c24e4e 100644 --- a/s3api/utils/unsigned-chunk-reader.go +++ b/s3api/utils/unsigned-chunk-reader.go @@ -39,22 +39,64 @@ var ( minChunkSize int64 = 8192 ) +// UnsignedChunkReader decodes AWS aws-chunked unsigned streaming request +// bodies. It strips chunk headers/trailers, validates chunk framing and trailing +// checksums, and exposes only object payload bytes through Read. +// +// The reader is intentionally streaming: chunk payload bytes are read directly +// into the caller's buffer. It only keeps bounded header/trailer parsing state +// and a few counters, so a large client-declared chunk does not become a large +// allocation in the gateway. type UnsignedChunkReader struct { reader *bufio.Reader checksumType checksumType parsedChecksum string hasher hash.Hash - stash []byte - offset int - // this data is necessary for 'InvalidChunkSizeError' error - // TODO: add 'Chunk' and 'BadChunkSize' in the error - chunkSizes []int64 - cLength int64 + // chunkDataLeft is the number of object data bytes still unread from the + // current chunk. These bytes are streamed directly into the caller's buffer + chunkDataLeft int64 + // needChunkEnd means the current chunk payload has been fully returned and + // the next Read must consume the chunk's trailing "\r\n" before parsing the + // next chunk header + needChunkEnd bool + // isEOF is set after the zero-sized chunk and trailer are parsed, so later + // reads return io.EOF without touching the underlying request body again. + isEOF bool + // The chunk-size rule needs information about the previous chunk: if the + // next parsed chunk is non-zero, the previous chunk was not the final chunk + // and must have been at least minChunkSize + chunkNumber int64 + lastChunkNumber int64 + lastChunkSize int64 + seenChunk bool + // TODO: Keep these fields ready for the future InvalidChunkSizeError shape: + // should be invalidChunkNumber and should be + // invalidChunkSize + invalidChunkNumber int64 + invalidChunkSize int64 + + cLength int64 // This data is necessary for the decoded content length mismatch error // TODO: add 'NumberBytesExpected' and 'NumberBytesProvided' in the error dataRead int64 } +func (ucr *UnsignedChunkReader) decodedBytesReturned() int64 { + return ucr.dataRead - ucr.chunkDataLeft +} + +func (ucr *UnsignedChunkReader) logState(prefix string) { + debuglogger.Logf("%s:\n returned_decoded_bytes=%v\n declared_chunk_bytes=%v\n current_chunk_left=%v\n chunk_number=%v\n last_chunk_number=%v\n last_chunk_size=%v\n need_chunk_end=%v", + prefix, + ucr.decodedBytesReturned(), + ucr.dataRead, + ucr.chunkDataLeft, + ucr.chunkNumber, + ucr.lastChunkNumber, + ucr.lastChunkSize, + ucr.needChunkEnd) +} + func NewUnsignedChunkReader(r io.Reader, ct checksumType, decContentLength int64) (*UnsignedChunkReader, error) { var hasher hash.Hash var err error @@ -62,17 +104,15 @@ func NewUnsignedChunkReader(r io.Reader, ct checksumType, decContentLength int64 hasher, err = getHasher(ct) } if err != nil { - debuglogger.Logf("failed to initialize hash calculator: %v", err) + debuglogger.Logf("unsigned chunk reader failed to initialize hash calculator for trailing checksum type %q and decoded content length %v: %v", ct, decContentLength, err) return nil, err } - debuglogger.Infof("initializing unsigned chunk reader") + debuglogger.Infof("initializing unsigned chunk reader:\n decoded_content_length=%v\n checksum_type=%q", decContentLength, ct) return &UnsignedChunkReader{ - reader: bufio.NewReader(r), + reader: bufio.NewReaderSize(r, maxHeaderSize), checksumType: ct, - stash: make([]byte, 0), hasher: hasher, - chunkSizes: []int64{}, cLength: decContentLength, }, nil } @@ -88,93 +128,170 @@ func (ucr *UnsignedChunkReader) Checksum() string { } func (ucr *UnsignedChunkReader) Read(p []byte) (int, error) { - // First read any stashed data - if len(ucr.stash) != 0 { - debuglogger.Infof("recovering the stash: (stash length): %v", len(ucr.stash)) - n := copy(p, ucr.stash) - ucr.offset += n + if len(p) == 0 { + return 0, nil + } - if n < len(ucr.stash) { - ucr.stash = ucr.stash[n:] - ucr.offset = 0 + if ucr.isEOF { + return 0, io.EOF + } + + var n int + + for n < len(p) { + // Once a chunk body is drained, validate its CRLF boundary before any + // more data can be returned. This preserves chunk framing validation + // while still allowing the body bytes themselves to pass through. + if ucr.needChunkEnd { + if err := ucr.readAndSkip('\r', '\n'); err != nil { + debuglogger.Logf("unsigned chunk reader failed to validate chunk payload delimiter: expected trailing \\r\\n after chunk %v with %v bytes already copied into caller buffer: %v", ucr.lastChunkNumber, n, err) + ucr.logState("unsigned chunk reader state after chunk payload delimiter failure") + return n, err + } + ucr.needChunkEnd = false + } + + if ucr.chunkDataLeft == 0 { + // No payload is pending, so the next bytes must be the chunk-size + chunkSize, err := ucr.extractChunkSize() + if err != nil { + debuglogger.Logf("unsigned chunk reader failed to parse next chunk header after copying %v bytes into caller buffer: %v", n, err) + ucr.logState("unsigned chunk reader state after chunk header parse failure") + return n, err + } + + if chunkSize == 0 { + // The zero-sized chunk ends the object data stream. At this + // point all declared chunk payload bytes have been consumed, so + // validate the decoded content length and then parse trailers. + ucr.isEOF = true + if ucr.cLength != ucr.dataRead { + debuglogger.Logf("unsigned chunk reader decoded content length mismatch at final chunk: expected %v decoded bytes, parsed %v decoded bytes from chunk headers", ucr.cLength, ucr.dataRead) + ucr.logState("unsigned chunk reader state after decoded content length mismatch") + err := s3err.GetAPIError(s3err.ErrContentLengthMismatch) + return n, err + } + + if err := ucr.readTrailer(); err != nil { + debuglogger.Logf("unsigned chunk reader failed to parse or validate trailers after final chunk: %v", err) + ucr.logState("unsigned chunk reader state after trailer failure") + return n, err + } + + return n, io.EOF + } + + ucr.dataRead += chunkSize + ucr.chunkDataLeft = chunkSize + } + + contentLeft := ucr.remainingContentLength() + if contentLeft == 0 && ucr.chunkDataLeft > 0 { + // The client declared more chunk payload bytes than the decoded + // content length allows. Do not pass those bytes to the backend + // writer; return the S3 error from this reader instead. + return n, ucr.handleExcessChunkData() + } + + // Read only as much object data as fits in p, the current chunk, and the + // decoded content length. This is the key streaming path: data is copied + // from the request body into p without allocating a chunk-sized buffer. + limit := min(int64(len(p)-n), ucr.chunkDataLeft, contentLeft) + readEnd := int64(n) + limit + read, err := ucr.reader.Read(p[n:readEnd]) + if read > 0 { + if ucr.hasher != nil { + if _, hashErr := ucr.hasher.Write(p[n : n+read]); hashErr != nil { + debuglogger.Logf("unsigned chunk reader failed to update trailing checksum hash after reading %v bytes from chunk %v: %v", read, ucr.lastChunkNumber, hashErr) + ucr.logState("unsigned chunk reader state after checksum hash failure") + return n, hashErr + } + } + ucr.chunkDataLeft -= int64(read) + n += read + if ucr.chunkDataLeft == 0 { + ucr.needChunkEnd = true + } + } + if err != nil { + if errors.Is(err, io.EOF) { + debuglogger.Logf("unsigned chunk reader reached EOF while reading chunk %v payload: copied %v bytes into caller buffer, %v bytes still expected in current chunk", ucr.lastChunkNumber, n, ucr.chunkDataLeft) + err = s3err.GetAPIError(s3err.ErrIncompleteBody) + } else { + debuglogger.Logf("unsigned chunk reader failed while reading chunk %v payload after copying %v bytes into caller buffer: %v", ucr.lastChunkNumber, n, err) + } + ucr.logState("unsigned chunk reader state after chunk payload read failure") + return n, err + } + if read == 0 { + debuglogger.Infof("unsigned chunk reader read zero bytes from underlying reader before filling caller buffer:\n buffer_size=%v\n returned_bytes=%v", len(p), n) return n, nil } } - for { - // Read the chunk size - chunkSize, err := ucr.extractChunkSize() - if err != nil { - return 0, err - } + debuglogger.Infof("Read:\n buffer_size=%v\n bytes_returned_this_read=%v\n decoded_bytes_returned_total=%v\n current_chunk_left_for_next_read=%v\n need_chunk_end_for_next_read=%v\n current_chunk_number=%v\n declared_chunk_bytes=%v", + len(p), + n, + ucr.decodedBytesReturned(), + ucr.chunkDataLeft, + ucr.needChunkEnd, + ucr.lastChunkNumber, + ucr.dataRead) - ucr.dataRead += chunkSize + return n, nil +} - if chunkSize == 0 { - // Stop reading parsing payloads as 0 sized chunk is reached - break - } - var rdr io.Reader = ucr.reader - if ucr.hasher != nil { - rdr = io.TeeReader(ucr.reader, ucr.hasher) - } - payload := make([]byte, chunkSize) - // Read and cache the payload - _, err = io.ReadFull(rdr, payload) - if err != nil { - // the chunk size is not 0 and if io.EOF is returned - // it means the body is incomplete - if errors.Is(err, io.EOF) { - debuglogger.Logf("unexpected EOF when reading chunk data") - return 0, s3err.GetAPIError(s3err.ErrIncompleteBody) - } - debuglogger.Logf("failed to read chunk data: %v", err) - return 0, err - } - - // Skip the trailing "\r\n" - if err := ucr.readAndSkip('\r', '\n'); err != nil { - debuglogger.Logf("failed to read trailing \\r\\n after chunk data: %v", err) - return 0, err - } - - // Copy the payload into the io.Reader buffer - n := copy(p[ucr.offset:], payload) - ucr.offset += n - - if int64(n) < chunkSize { - // stash the remaining data - ucr.stash = payload[n:] - debuglogger.Infof("stashing the remaining data: (stash length): %v", len(ucr.stash)) - dataRead := ucr.offset - ucr.offset = 0 - return dataRead, nil - } +func (ucr *UnsignedChunkReader) remainingContentLength() int64 { + // dataRead is the sum of parsed chunk sizes, while chunkDataLeft is the + // unread part of the current chunk. Their difference is the decoded object + // byte count already returned or ready to return to the caller. + read := ucr.dataRead - ucr.chunkDataLeft + if read >= ucr.cLength { + return 0 } - if ucr.cLength != ucr.dataRead { - debuglogger.Logf("number of bytes expected: (%v), number of bytes read: (%v)", ucr.cLength, ucr.dataRead) - return 0, s3err.GetAPIError(s3err.ErrContentLengthMismatch) + return ucr.cLength - read +} + +func (ucr *UnsignedChunkReader) handleExcessChunkData() error { + // When the decoded content length is exhausted in the middle of a chunk, + // distinguish "extra payload data" from "chunk ended before its declared + // size". The latter is an incomplete body; the former is a content-length + // mismatch. Peek keeps the bytes buffered and avoids forwarding either case + // to the backend writer. + buf, err := ucr.reader.Peek(2) + if len(buf) > 0 && buf[0] != '\r' { + debuglogger.Logf("unsigned chunk reader decoded content length exhausted in chunk %v, but next byte is payload data instead of chunk delimiter: expected decoded length %v", ucr.lastChunkNumber, ucr.cLength) + ucr.logState("unsigned chunk reader state after excess chunk payload detection") + return s3err.GetAPIError(s3err.ErrContentLengthMismatch) + } + if len(buf) > 1 && buf[1] != '\n' { + debuglogger.Logf("unsigned chunk reader decoded content length exhausted in chunk %v, but next bytes are not a valid chunk delimiter: got %q, expected \\r\\n", ucr.lastChunkNumber, buf) + ucr.logState("unsigned chunk reader state after invalid delimiter at decoded length boundary") + return s3err.GetAPIError(s3err.ErrContentLengthMismatch) + } + if err != nil { + debuglogger.Logf("unsigned chunk reader could not peek chunk delimiter after decoded content length was exhausted: %v", err) + ucr.logState("unsigned chunk reader state after delimiter peek failure") + return s3err.GetAPIError(s3err.ErrIncompleteBody) } - // Read and validate trailers - if err := ucr.readTrailer(); err != nil { - debuglogger.Logf("failed to read trailer: %v", err) - return 0, err - } - - return ucr.offset, io.EOF + debuglogger.Logf("unsigned chunk reader found a chunk delimiter before declared chunk %v payload was fully read: %v bytes still expected", ucr.lastChunkNumber, ucr.chunkDataLeft) + ucr.logState("unsigned chunk reader state after short chunk payload detection") + return s3err.GetAPIError(s3err.ErrIncompleteBody) } // Reads and validates the bytes provided from the underlying io.Reader func (ucr *UnsignedChunkReader) readAndSkip(data ...byte) error { - for _, d := range data { + for i, d := range data { b, err := ucr.reader.ReadByte() if err != nil { + debuglogger.Logf("unsigned chunk reader failed to read expected byte %d of delimiter %q: expected %q, err: %v", i+1, data, d, err) return s3err.GetAPIError(s3err.ErrIncompleteBody) } if b != d { + debuglogger.Logf("unsigned chunk reader delimiter mismatch at byte %d of %q: expected %q, got %q", i+1, data, d, b) return s3err.GetAPIError(s3err.ErrIncompleteBody) } } @@ -184,50 +301,78 @@ func (ucr *UnsignedChunkReader) readAndSkip(data ...byte) error { // Extracts the chunk size from the payload func (ucr *UnsignedChunkReader) extractChunkSize() (int64, error) { - line, err := ucr.reader.ReadString('\r') + line, err := ucr.readChunkSizeLine() if err != nil { - debuglogger.Logf("failed to parse chunk size: %v", err) + debuglogger.Logf("unsigned chunk reader failed to read chunk size line for chunk %v: %v", ucr.chunkNumber+1, err) return 0, s3err.GetAPIError(s3err.ErrIncompleteBody) } - err = ucr.readAndSkip('\n') - if err != nil { - debuglogger.Logf("failed to read the second byte (\\n) after chunk size") - return 0, err - } - - line = strings.TrimSpace(line) - chunkSize, err := strconv.ParseInt(line, 16, 64) - if err != nil { - debuglogger.Logf("failed to convert chunk size: %v", err) + if err != nil || chunkSize < 0 { + debuglogger.Logf("unsigned chunk reader failed to parse chunk %v size %q as non-negative hexadecimal int64: %v", ucr.chunkNumber+1, line, err) return 0, s3err.GetAPIError(s3err.ErrIncompleteBody) } + ucr.chunkNumber++ if !ucr.isValidChunkSize(chunkSize) { + debuglogger.Logf("unsigned chunk reader invalid chunk size detected while parsing chunk %v: previous chunk %v had size %v, current chunk size is %v", ucr.chunkNumber, ucr.invalidChunkNumber, ucr.invalidChunkSize, chunkSize) return chunkSize, s3err.GetAPIError(s3err.ErrInvalidChunkSize) } - ucr.chunkSizes = append(ucr.chunkSizes, chunkSize) + ucr.lastChunkNumber = ucr.chunkNumber + ucr.lastChunkSize = chunkSize + ucr.seenChunk = true debuglogger.Infof("chunk size extracted: %v", chunkSize) return chunkSize, nil } +func (ucr *UnsignedChunkReader) readChunkSizeLine() (string, error) { + var line []byte + for { + // ReadSlice lets normal headers use bufio's internal buffer. The append + // path only handles split or oversized headers and is bounded by + // maxHeaderSize, so malformed headers cannot grow memory unboundedly. + part, err := ucr.reader.ReadSlice('\r') + line = append(line, part...) + if len(line) > maxHeaderSize { + debuglogger.Logf("unsigned chunk reader chunk %v size header exceeds maximum allowed size: header_len=%v, header_limit=%v", ucr.chunkNumber+1, len(line), maxHeaderSize) + return "", s3err.GetAPIError(s3err.ErrIncompleteBody) + } + if err == nil { + break + } + if errors.Is(err, bufio.ErrBufferFull) { + continue + } + debuglogger.Logf("unsigned chunk reader failed while reading chunk %v size header before \\r delimiter: %v", ucr.chunkNumber+1, err) + return "", err + } + + err := ucr.readAndSkip('\n') + if err != nil { + debuglogger.Logf("unsigned chunk reader chunk %v size header is not followed by \\n after \\r: %v", ucr.chunkNumber+1, err) + return "", err + } + + return strings.TrimSpace(string(line)), nil +} + // isValidChunkSize checks if the parsed chunk size is valid // they follow one rule: all chunk sizes except for the last one // should be greater than 8192 func (ucr *UnsignedChunkReader) isValidChunkSize(size int64) bool { - if len(ucr.chunkSizes) == 0 { + if !ucr.seenChunk { // any valid number is valid as a first chunk size return true } - lastChunkSize := ucr.chunkSizes[len(ucr.chunkSizes)-1] // any chunk size, except the last one should be greater than 8192 - if size != 0 && lastChunkSize < minChunkSize { - debuglogger.Logf("invalid chunk size %v", lastChunkSize) + if size != 0 && ucr.lastChunkSize < minChunkSize { + ucr.invalidChunkNumber = ucr.lastChunkNumber + ucr.invalidChunkSize = ucr.lastChunkSize + debuglogger.Logf("unsigned chunk reader previous chunk is too small to be followed by another data chunk:\n invalid_chunk_number=%v\n bad_chunk_size=%v\n min_chunk_size=%v\n next_chunk_size=%v", ucr.invalidChunkNumber, ucr.invalidChunkSize, minChunkSize, size) return false } @@ -242,6 +387,7 @@ func (ucr *UnsignedChunkReader) readTrailer() error { for { v, err := ucr.reader.ReadByte() if err != nil { + debuglogger.Logf("unsigned chunk reader failed to read trailer byte after final chunk: %v", err) return s3err.GetAPIError(s3err.ErrIncompleteBody) } if v != '\r' { @@ -256,7 +402,7 @@ func (ucr *UnsignedChunkReader) readTrailer() error { // only read the last byte: \n err := ucr.readAndSkip('\n') if err != nil { - debuglogger.Logf("failed to read chunk last byte: \\n: %v", err) + debuglogger.Logf("unsigned chunk reader empty trailer terminator is incomplete: expected final \\n after \\r: %v", err) return s3err.GetAPIError(s3err.ErrIncompleteBody) } @@ -266,11 +412,11 @@ func (ucr *UnsignedChunkReader) readTrailer() error { var tmp [3]byte _, err = io.ReadFull(ucr.reader, tmp[:]) if err != nil { - debuglogger.Logf("failed to read chunk ending: \\n\\r\\n: %v", err) + debuglogger.Logf("unsigned chunk reader trailer delimiter is incomplete after trailer header %q: expected \\n\\r\\n, err: %v", trailerBuffer.String(), err) return s3err.GetAPIError(s3err.ErrIncompleteBody) } if !bytes.Equal(tmp[:], trailerDelim) { - debuglogger.Logf("incorrect trailer delimiter: (expected): \\n\\r\\n, (got): %q", tmp[:]) + debuglogger.Logf("unsigned chunk reader trailer delimiter mismatch after trailer header %q: expected \\n\\r\\n, got %q", trailerBuffer.String(), tmp[:]) return s3err.GetAPIError(s3err.ErrIncompleteBody) } break @@ -281,7 +427,7 @@ func (ucr *UnsignedChunkReader) readTrailer() error { trailerHeader = strings.TrimSpace(trailerHeader) if trailerHeader == "" { if ucr.checksumType != "" { - debuglogger.Logf("expected %s checksum in the paylod, but it's missing", ucr.checksumType) + debuglogger.Logf("unsigned chunk reader expected trailing checksum %s, but final trailer header is empty", ucr.checksumType) return s3err.GetAPIError(s3err.ErrMalformedTrailer) } @@ -289,7 +435,7 @@ func (ucr *UnsignedChunkReader) readTrailer() error { } trailerHeaderParts := strings.Split(trailerHeader, ":") if len(trailerHeaderParts) != 2 { - debuglogger.Logf("invalid trailer header parts: %v", trailerHeaderParts) + debuglogger.Logf("unsigned chunk reader malformed trailer header %q: expected exactly one ':' separator, got %v parts", trailerHeader, len(trailerHeaderParts)) return s3err.GetAPIError(s3err.ErrMalformedTrailer) } @@ -297,17 +443,17 @@ func (ucr *UnsignedChunkReader) readTrailer() error { checksum := trailerHeaderParts[1] if !checksumKey.isValid() { - debuglogger.Logf("invalid checksum header key: %s", checksumKey) + debuglogger.Logf("unsigned chunk reader malformed trailer header %q: unsupported checksum key %q", trailerHeader, checksumKey) return s3err.GetAPIError(s3err.ErrMalformedTrailer) } if checksumKey != ucr.checksumType { - debuglogger.Logf("incorrect checksum type (expected): %s, (actual): %s", ucr.checksumType, checksumKey) + debuglogger.Logf("unsigned chunk reader trailer checksum type mismatch: expected %q from x-amz-trailer, got %q in trailer header", ucr.checksumType, checksumKey) return s3err.GetAPIError(s3err.ErrMalformedTrailer) } ucr.parsedChecksum = checksum - debuglogger.Infof("parsed the trailing header:\n%v:%v", checksumKey, checksum) + debuglogger.Infof("parsed the trailing header:%s:%s", checksumKey, checksum) // Validate checksum return ucr.validateChecksum() @@ -318,7 +464,7 @@ func (ucr *UnsignedChunkReader) validateChecksum() error { algo := types.ChecksumAlgorithm(strings.ToUpper(strings.TrimPrefix(string(ucr.checksumType), "x-amz-checksum-"))) // validate the checksum if !IsValidChecksum(ucr.parsedChecksum, algo) { - debuglogger.Logf("invalid checksum: (algo): %s, (checksum): %s", algo, ucr.parsedChecksum) + debuglogger.Logf("unsigned chunk reader parsed trailing checksum has invalid format: algo=%s, checksum=%s", algo, ucr.parsedChecksum) return s3err.GetInvalidTrailingChecksumHeaderErr(string(ucr.checksumType)) } @@ -326,7 +472,7 @@ func (ucr *UnsignedChunkReader) validateChecksum() error { // compare the calculated and parsed checksums if checksum != ucr.parsedChecksum { - debuglogger.Logf("incorrect checksum: (expected): %v, (got): %v", ucr.parsedChecksum, checksum) + debuglogger.Logf("unsigned chunk reader trailing checksum mismatch: algo=%s, parsed_checksum=%v, calculated_checksum=%v, decoded_bytes=%v", algo, ucr.parsedChecksum, checksum, ucr.decodedBytesReturned()) return s3err.GetChecksumBadDigestErr(algo) } @@ -339,7 +485,7 @@ func (ucr *UnsignedChunkReader) calculateChecksum() string { return base64.StdEncoding.EncodeToString(csum) } -// Retruns the hash calculator based on the hash type provided +// Returns the hash calculator based on the hash type provided func getHasher(ct checksumType) (hash.Hash, error) { switch ct { case checksumTypeCrc32: diff --git a/s3api/utils/unsigned-chunk-reader_test.go b/s3api/utils/unsigned-chunk-reader_test.go new file mode 100644 index 00000000..d3eec17b --- /dev/null +++ b/s3api/utils/unsigned-chunk-reader_test.go @@ -0,0 +1,177 @@ +// Copyright 2026 Versity Software +// This file is licensed under the Apache License, Version 2.0 +// (the "License"); you may not use this file except in compliance +// with the License. You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, +// software distributed under the License is distributed on an +// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +// KIND, either express or implied. See the License for the +// specific language governing permissions and limitations +// under the License. + +package utils + +import ( + "bytes" + "crypto/sha256" + "encoding/base64" + "errors" + "fmt" + "io" + "strings" + "testing" + + "github.com/versity/versitygw/s3err" +) + +func TestUnsignedChunkReaderStreamsLargeChunkWithoutBuffering(t *testing.T) { + const chunkSize int64 = 1 << 32 + body := io.MultiReader( + strings.NewReader(fmt.Sprintf("%x\r\n", chunkSize)), + strings.NewReader("abc"), + ) + reader, err := NewUnsignedChunkReader(body, "", chunkSize) + if err != nil { + t.Fatal(err) + } + + buf := make([]byte, 3) + n, err := reader.Read(buf) + if err != nil { + t.Fatalf("read: %v", err) + } + if got := string(buf[:n]); got != "abc" { + t.Fatalf("read data = %q, want %q", got, "abc") + } + if reader.chunkDataLeft != chunkSize-int64(n) { + t.Fatalf("chunkDataLeft = %d, want %d", reader.chunkDataLeft, chunkSize-int64(n)) + } +} + +func TestUnsignedChunkReaderReadsAcrossChunksAndThenEOF(t *testing.T) { + firstChunk := strings.Repeat("a", int(minChunkSize)) + body := unsignedChunkBody(firstChunk, "tail") + reader, err := NewUnsignedChunkReader(strings.NewReader(body), "", int64(len(firstChunk)+len("tail"))) + if err != nil { + t.Fatal(err) + } + + var out bytes.Buffer + buf := make([]byte, 3) + for { + n, err := reader.Read(buf) + out.Write(buf[:n]) + if err == io.EOF { + break + } + if err != nil { + t.Fatalf("read: %v", err) + } + } + + expected := firstChunk + "tail" + if got := out.String(); got != expected { + t.Fatalf("read data length = %d, want %d", len(got), len(expected)) + } + + n, err := reader.Read(buf) + if n != 0 || err != io.EOF { + t.Fatalf("second EOF read = (%d, %v), want (0, EOF)", n, err) + } +} + +func TestUnsignedChunkReaderValidatesTrailingChecksum(t *testing.T) { + payload := "abcdefg" + sum := sha256.Sum256([]byte(payload)) + checksum := base64.StdEncoding.EncodeToString(sum[:]) + body := fmt.Sprintf("%x\r\n%s\r\n0\r\n%s:%s\r\n\r\n", + len(payload), payload, checksumTypeSha256, checksum) + + reader, err := NewUnsignedChunkReader(strings.NewReader(body), checksumTypeSha256, int64(len(payload))) + if err != nil { + t.Fatal(err) + } + out, err := io.ReadAll(reader) + if err != nil { + t.Fatalf("read all: %v", err) + } + if string(out) != payload { + t.Fatalf("read data = %q, want %q", out, payload) + } + if reader.Checksum() != checksum { + t.Fatalf("checksum = %q, want %q", reader.Checksum(), checksum) + } +} + +func TestUnsignedChunkReaderContentLengthMismatchStopsAtDecodedLength(t *testing.T) { + body := "b\r\nabcdefghijk\r\n0\r\n\r\n" + reader, err := NewUnsignedChunkReader(strings.NewReader(body), "", 5) + if err != nil { + t.Fatal(err) + } + + buf := make([]byte, 32) + n, err := reader.Read(buf) + if string(buf[:n]) != "abcde" { + t.Fatalf("read data = %q, want %q", buf[:n], "abcde") + } + requireAPIErrorCode(t, err, s3err.GetAPIError(s3err.ErrContentLengthMismatch).Code) +} + +func TestUnsignedChunkReaderDeclaredChunkLongerThanPayloadReturnsIncompleteBody(t *testing.T) { + body := "B\r\ndummy data\r\n0\r\n\r\n" + reader, err := NewUnsignedChunkReader(strings.NewReader(body), checksumTypeCrc64nvme, 10) + if err != nil { + t.Fatal(err) + } + + buf := make([]byte, 32) + n, err := reader.Read(buf) + if string(buf[:n]) != "dummy data" { + t.Fatalf("read data = %q, want %q", buf[:n], "dummy data") + } + requireAPIErrorCode(t, err, s3err.GetAPIError(s3err.ErrIncompleteBody).Code) +} + +func TestUnsignedChunkReaderInvalidChunkSize(t *testing.T) { + body := unsignedChunkBody("short", "x") + reader, err := NewUnsignedChunkReader(strings.NewReader(body), "", int64(len("short")+len("x"))) + if err != nil { + t.Fatal(err) + } + + _, err = io.ReadAll(reader) + requireAPIErrorCode(t, err, s3err.GetAPIError(s3err.ErrInvalidChunkSize).Code) + if reader.invalidChunkNumber != 1 { + t.Fatalf("invalidChunkNumber = %d, want 1", reader.invalidChunkNumber) + } + if reader.invalidChunkSize != int64(len("short")) { + t.Fatalf("invalidChunkSize = %d, want %d", reader.invalidChunkSize, len("short")) + } +} + +func unsignedChunkBody(chunks ...string) string { + var b strings.Builder + for _, chunk := range chunks { + fmt.Fprintf(&b, "%x\r\n%s\r\n", len(chunk), chunk) + } + b.WriteString("0\r\n\r\n") + return b.String() +} + +func requireAPIErrorCode(t *testing.T, err error, code string) { + t.Helper() + if err == nil { + t.Fatalf("expected %s error, got nil", code) + } + var apiErr s3err.APIError + if !errors.As(err, &apiErr) { + t.Fatalf("expected APIError, got %T: %v", err, err) + } + if apiErr.Code != code { + t.Fatalf("APIError code = %q, want %q", apiErr.Code, code) + } +}