diff --git a/weed/server/common.go b/weed/server/common.go index f719d6881..b45fe3a5c 100644 --- a/weed/server/common.go +++ b/weed/server/common.go @@ -288,12 +288,32 @@ func adjustHeaderContentDisposition(w http.ResponseWriter, r *http.Request, file } } +// responseBodyTracker records how many bytes have reached the ResponseWriter. +// Bytes still sitting in the pooled bufio.Writer do not count: the status line +// is not committed until the buffer flushes. +type responseBodyTracker struct { + http.ResponseWriter + written int64 +} + +func (t *responseBodyTracker) Write(p []byte) (int, error) { + n, err := t.ResponseWriter.Write(p) + t.written += int64(n) + return n, err +} + func ProcessRangeRequest(r *http.Request, w http.ResponseWriter, totalSize int64, mimeType string, prepareWriteFn func(offset int64, size int64) (filer.DoStreamContent, error)) error { rangeReq := r.Header.Get("Range") + sink := &responseBodyTracker{ResponseWriter: w} bufferedWriter := writePool.Get().(*bufio.Writer) - bufferedWriter.Reset(w) + bufferedWriter.Reset(sink) + discardBuffered := false defer func() { - bufferedWriter.Flush() + if discardBuffered { + bufferedWriter.Reset(io.Discard) + } else { + bufferedWriter.Flush() + } writePool.Put(bufferedWriter) }() @@ -306,6 +326,13 @@ func ProcessRangeRequest(r *http.Request, w http.ResponseWriter, totalSize int64 } if err = writeFn(bufferedWriter); err != nil { glog.Errorf("ProcessRangeRequest: %v", err) + // Drop the unflushed tail: flushing it would finish a 200 whose + // Content-Length is already set, delivering corrupt bytes as a + // complete body. + discardBuffered = true + if sink.written > 0 { + panic(http.ErrAbortHandler) + } w.Header().Del("Content-Length") http.Error(w, err.Error(), http.StatusInternalServerError) return fmt.Errorf("ProcessRangeRequest: %w", err) @@ -359,8 +386,10 @@ func ProcessRangeRequest(r *http.Request, w http.ResponseWriter, totalSize int64 err = writeFn(bufferedWriter) if err != nil { glog.Errorf("ProcessRangeRequest range[0]: %+v err: %v", w.Header(), err) - // Cannot call http.Error() here because WriteHeader was already called - return fmt.Errorf("ProcessRangeRequest range[0]: %w", err) + // WriteHeader was already called: drop the unflushed tail and + // abort so the client does not read corrupt bytes as a full body. + discardBuffered = true + panic(http.ErrAbortHandler) } return nil } @@ -412,8 +441,9 @@ func ProcessRangeRequest(r *http.Request, w http.ResponseWriter, totalSize int64 w.WriteHeader(http.StatusPartialContent) if _, err := io.CopyN(bufferedWriter, sendContent, sendSize); err != nil { glog.Errorf("ProcessRangeRequest err: %v", err) - // Cannot call http.Error() here because WriteHeader was already called - return fmt.Errorf("ProcessRangeRequest err: %w", err) + // WriteHeader was already called: drop the unflushed tail and abort. + discardBuffered = true + panic(http.ErrAbortHandler) } return nil } diff --git a/weed/server/common_test.go b/weed/server/common_test.go index 8042f78d7..ea81fabca 100644 --- a/weed/server/common_test.go +++ b/weed/server/common_test.go @@ -3,6 +3,7 @@ package weed_server import ( "bytes" "context" + "errors" "io" "mime/multipart" "net/http" @@ -219,6 +220,85 @@ func (c *countingReadCloser) Close() error { return nil } +// A write error after the body has flushed must not complete as 200 with the +// full Content-Length: the unflushed tail is discarded and the connection +// aborted mid-body. +func TestProcessRangeRequestWriteErrorAfterBodyStarted(t *testing.T) { + payload := bytes.Repeat([]byte("x"), 200*1024) + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _ = ProcessRangeRequest(r, w, int64(len(payload)), "application/octet-stream", func(offset int64, size int64) (filer.DoStreamContent, error) { + return func(writer io.Writer) error { + // Page-sized writes, as readNeedleDataInto does. The tail sits + // in the 128KiB buffer; flushing it would finish the 200. + for off := 0; off < len(payload); off += 4096 { + end := off + 4096 + if end > len(payload) { + end = len(payload) + } + if _, werr := writer.Write(payload[off:end]); werr != nil { + return werr + } + } + return errors.New("ReadNeedleData checksum mismatch") + }, nil + }) + })) + defer srv.Close() + + resp, err := srv.Client().Get(srv.URL) + if err != nil { + t.Fatalf("request: %v", err) + } + defer resp.Body.Close() + if resp.StatusCode != http.StatusOK { + t.Fatalf("status %d, want the committed 200", resp.StatusCode) + } + body, readErr := io.ReadAll(resp.Body) + if readErr == nil { + t.Fatalf("body completed with %d bytes; the erroring tail must abort it", len(body)) + } +} + +func TestProcessRangeRequestWriteErrorBeforeBody(t *testing.T) { + r := httptest.NewRequest(http.MethodGet, "/test.bin", nil) + w := httptest.NewRecorder() + err := ProcessRangeRequest(r, w, 100, "application/octet-stream", func(offset int64, size int64) (filer.DoStreamContent, error) { + return func(writer io.Writer) error { + return errors.New("ReadNeedleData checksum mismatch") + }, nil + }) + if err == nil { + t.Fatal("expected write error") + } + if w.Code != http.StatusInternalServerError { + t.Fatalf("status %d, want 500, body %q", w.Code, w.Body.String()) + } + if !strings.Contains(w.Body.String(), "checksum") { + t.Fatalf("body %q, want the checksum error", w.Body.String()) + } +} + +func TestProcessRangeRequestLargeBody(t *testing.T) { + payload := bytes.Repeat([]byte("y"), 200*1024) + r := httptest.NewRequest(http.MethodGet, "/test.bin", nil) + w := httptest.NewRecorder() + err := ProcessRangeRequest(r, w, int64(len(payload)), "application/octet-stream", func(offset int64, size int64) (filer.DoStreamContent, error) { + return func(writer io.Writer) error { + _, werr := writer.Write(payload) + return werr + }, nil + }) + if err != nil { + t.Fatalf("write: %v", err) + } + if w.Code != http.StatusOK { + t.Fatalf("status %d, want 200", w.Code) + } + if !bytes.Equal(w.Body.Bytes(), payload) { + t.Fatalf("body len %d, want %d", w.Body.Len(), len(payload)) + } +} + func TestProcessRangeRequestRanges(t *testing.T) { data := []byte("0123456789") serve := func(rangeHeader string) (*httptest.ResponseRecorder, error) { diff --git a/weed/storage/volume_read.go b/weed/storage/volume_read.go index 91292c4ed..672d774f6 100644 --- a/weed/storage/volume_read.go +++ b/weed/storage/volume_read.go @@ -163,10 +163,21 @@ func (v *Volume) readNeedleDataInto(n *needle.Needle, readOption *ReadOption, wr } buf := mem.Allocate(min(readOption.ReadBufferSize, int(size))) - defer mem.Free(buf) + // A full-needle read holds the last page back in `pending` until the CRC + // matches; `spare` is a second buffer swapped in so the held page is not + // overwritten while streaming. + var spare []byte + defer func() { + mem.Free(buf) + if spare != nil { + mem.Free(spare) + } + }() // read needle data crc := needle.CRC(0) + var pending []byte + checkCRC := offset == 0 && size == int64(n.DataSize) for x := offset; x < offset+size; x += int64(len(buf)) { if readOption.HasSlowRead { @@ -212,9 +223,22 @@ func (v *Volume) readNeedleDataInto(n *needle.Needle, readOption *ReadOption, wr toWrite := min(count, int(offset+size-x)) if toWrite > 0 { crc = crc.Update(buf[0:toWrite]) - // Note: CRC validation happens after the loop completes (see below) - // to avoid performance overhead in the hot read path - if _, err = writer.Write(buf[0:toWrite]); err != nil { + // The CRC is known only after the last byte; hold each page until + // the next one is read so a bad needle is never fully written. + if checkCRC { + if pending != nil { + if _, err = writer.Write(pending); err != nil { + return fmt.Errorf("ReadNeedleData write: %w", err) + } + } + pending = buf[:toWrite] + if x+int64(len(buf)) < offset+size { + if spare == nil { + spare = mem.Allocate(len(buf)) + } + buf, spare = spare, buf + } + } else if _, err = writer.Write(buf[0:toWrite]); err != nil { return fmt.Errorf("ReadNeedleData write: %w", err) } } @@ -234,13 +258,19 @@ func (v *Volume) readNeedleDataInto(n *needle.Needle, readOption *ReadOption, wr // we still return that error to the caller, but the disk itself // produced clean bytes. v.checkReadWriteError(nil) - if offset == 0 && size == int64(n.DataSize) && (n.Checksum != crc && uint32(n.Checksum) != crc.Value()) { + if checkCRC && (n.Checksum != crc && uint32(n.Checksum) != crc.Value()) { // the crc.Value() function is to be deprecated. this double checking is for backward compatibility // with seaweed version using crc.Value() instead of uint32(crc), which appears in commit 056c480eb // and switch appeared in version 3.09. + // pending is the last page and is intentionally not written. stats.VolumeServerHandlerCounter.WithLabelValues(stats.ErrorCRC).Inc() return fmt.Errorf("ReadNeedleData checksum %v expected %v for Needle: %v,%v", crc, n.Checksum, v.Id, n) } + if pending != nil { + if _, err = writer.Write(pending); err != nil { + return fmt.Errorf("ReadNeedleData write: %w", err) + } + } return nil } diff --git a/weed/storage/volume_read_test.go b/weed/storage/volume_read_test.go index 9cc256c2d..be3971e74 100644 --- a/weed/storage/volume_read_test.go +++ b/weed/storage/volume_read_test.go @@ -7,6 +7,7 @@ import ( "os" "path/filepath" "reflect" + "strings" "testing" "time" @@ -242,3 +243,82 @@ func TestScanVolumeFileFrom_StopsAtRecordThatCannotAdvance(t *testing.T) { } } } + +// The CRC is known only after the last byte, so a full read must keep that +// page unwritten when the checksum mismatches; otherwise the GET has already +// committed a corrupt body. +func TestReadNeedleDataIntoChecksumMismatchHoldsLastPage(t *testing.T) { + dir := t.TempDir() + v, err := NewVolume(dir, dir, "", 1, NeedleMapInMemory, &super_block.ReplicaPlacement{}, &needle.TTL{}, 0, needle.GetCurrentVersion(), 0, 0) + if err != nil { + t.Fatalf("volume creation: %v", err) + } + defer v.Close() + + // 2048 is two pages (one buffer swap), 3000 is three, and 2500 with a + // buffer larger than the needle is the single Write that used to commit + // the whole body before the CRC check. + cases := []struct { + name string + size int + page int + }{ + {"two pages", 2048, 1024}, + {"three pages", 3000, 1024}, + {"one buffer", 2500, 4096}, + } + for i, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + data := bytes.Repeat([]byte("abcdefghij"), (tc.size+9)/10)[:tc.size] + n := new(needle.Needle) + n.Data = append([]byte(nil), data...) + n.Checksum = needle.NewCRC(n.Data) + n.Id = types.Uint64ToNeedleId(uint64(i + 1)) + offset, _, _, err := v.writeNeedle2(n, true, false, false) + if err != nil { + t.Fatalf("write needle: %v", err) + } + nv, ok := v.nm.Get(n.Id) + if !ok { + t.Fatal("needle missing from index") + } + actual := nv.Offset.ToActualOffset() + + read := func() (bytes.Buffer, error) { + t.Helper() + meta := new(needle.Needle) + meta.Id = n.Id + if err := v.readNeedleMetaAt(meta, actual, int32(nv.Size)); err != nil { + t.Fatalf("read meta at %d size %d: %v", actual, nv.Size, err) + } + var buf bytes.Buffer + err := v.readNeedleDataInto(meta, &ReadOption{ReadBufferSize: tc.page}, &buf, 0, int64(meta.DataSize)) + return buf, err + } + + intact, err := read() + if err != nil { + t.Fatalf("intact read: %v", err) + } + if !bytes.Equal(intact.Bytes(), data) { + t.Fatalf("intact read len %d, want %d", intact.Len(), len(data)) + } + + dataOff := int64(offset) + types.NeedleHeaderSize + types.DataSizeSize + if _, err := v.DataBackend.WriteAt([]byte{0xff}, dataOff); err != nil { + t.Fatalf("damage needle: %v", err) + } + damaged, err := read() + if err == nil || !strings.Contains(err.Error(), "checksum") { + t.Fatalf("damaged read: got %v, want a checksum error", err) + } + held := len(data) % tc.page + if held == 0 { + held = tc.page + } + if damaged.Len() != len(data)-held { + t.Fatalf("damaged read wrote %d bytes, want %d with the last page held back", damaged.Len(), len(data)-held) + } + }) + } +}