diff --git a/client_test.go b/client_test.go index a77b0a5..30b0733 100644 --- a/client_test.go +++ b/client_test.go @@ -1,6 +1,7 @@ package warc import ( + "bytes" "context" "crypto/rand" "crypto/rsa" @@ -1667,6 +1668,7 @@ func TestHTTPClientWithZStandard(t *testing.T) { for _, path := range files { testFileSingleHashCheck(t, path, "sha1:UIRWL5DFIPQ4MX3D3GFHM2HCVU3TZ6I3", []string{"26872"}, 1, server.URL+"/") + assertZSTDFileFramesValid(t, path) } } @@ -1713,6 +1715,43 @@ func TestHTTPClientWithZStandardDictionary(t *testing.T) { for _, path := range files { testFileSingleHashCheck(t, path, "sha1:UIRWL5DFIPQ4MX3D3GFHM2HCVU3TZ6I3", []string{"26872"}, 1, server.URL+"/") + assertZSTDFileFramesValid(t, path) + } +} + +func assertZSTDFileFramesValid(t *testing.T, path string) { + t.Helper() + + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("assertZSTDFileFramesValid: read %s: %v", path, err) + } + + // Skip any leading skippable frames which we add when we are embedding a dictionary. + for len(data) >= 8 { + m := uint32(data[0]) | uint32(data[1])<<8 | uint32(data[2])<<16 | uint32(data[3])<<24 + if m != 0x184D2A5D { + break + } + size := uint32(data[4]) | uint32(data[5])<<8 | uint32(data[6])<<16 | uint32(data[7])<<24 + data = data[8+size:] + } + + r := bytes.NewReader(data) + for i := 0; ; i++ { + frame, err := readZSTDFrameBytes(r) + if errors.Is(err, io.EOF) { + break + } + if err != nil { + t.Fatalf("assertZSTDFileFramesValid: read frame %d from %s: %v", i, path, err) + } + if !zstdFrameHasContentSize(frame) { + t.Errorf("ZSTD frame %d in %s is missing Frame_Content_Size", i, path) + } + if !zstdFrameHasChecksum(frame) { + t.Errorf("ZSTD frame %d in %s is missing Content_Checksum", i, path) + } } } diff --git a/dedupe.go b/dedupe.go index c9ea751..3c12ab4 100644 --- a/dedupe.go +++ b/dedupe.go @@ -37,7 +37,7 @@ type revisitRecord struct { responseUUID string targetURI string date time.Time - size int + size int64 } func (d *customDialer) checkLocalRevisit(digest string) revisitRecord { @@ -75,7 +75,7 @@ func checkCDXRevisit(CDXURL string, digest string, targetURI string, cookie stri cdxReply := strings.Fields(string(body)) if len(cdxReply) >= 7 && cdxReply[3] != "warc/revisit" && cdxReply[5] == digest { - recordSize, _ := strconv.Atoi(cdxReply[6]) + recordSize, _ := strconv.ParseInt(cdxReply[6], 10, 64) t, err := time.Parse("20060102150405", cdxReply[1]) if err != nil { diff --git a/dialer.go b/dialer.go index 9409ca0..ad6dae1 100644 --- a/dialer.go +++ b/dialer.go @@ -438,9 +438,9 @@ func (d *customDialer) CustomDialTLSContext(ctx context.Context, network, addres serverName := address if host, _, err := net.SplitHostPort(address); err != nil { - return nil, fmt.Errorf("failed to extract host from address %s: %w", address, err) + return nil, fmt.Errorf("failed to extract host from address %s: %w", address, err) } else { - serverName = host + serverName = host } cfg := &tls.Config{ @@ -611,7 +611,7 @@ func (d *customDialer) writeWARCFromConnection(ctx context.Context, reqPipe, res } r.Header.Set("WARC-Block-Digest", digest) - r.Header.Set("Content-Length", strconv.Itoa(getContentLength(r.Content))) + r.Header.Set("Content-Length", strconv.FormatInt(getContentLength(r.Content), 10)) if d.client.dedupeOptions.LocalDedupe { if r.Header.Get("WARC-Type") == "response" && !slices.Contains(emptyPayloadDigests, r.Header.Get("WARC-Payload-Digest")) { diff --git a/utils.go b/utils.go index 46274a1..0fca1db 100644 --- a/utils.go +++ b/utils.go @@ -91,14 +91,7 @@ func NewWriter(writer io.Writer, fileName string, digestAlgorithm DigestAlgorith } } - // Create ZStandard writer either with or without the encoder dictionary and return it. - var zstdWriter *zstd.Encoder - var err error - eopts := []zstd.EOption{zstd.WithEncoderLevel(zstd.SpeedBetterCompression)} - if len(dictionary) > 0 { - eopts = append(eopts, zstd.WithEncoderDict(dictionary)) - } - zstdWriter, err = zstd.NewWriter(writer, eopts...) + zstdWriter, err := newSizedZstdWriter(writer, dictionary) if err != nil { return nil, err } @@ -208,7 +201,7 @@ func checkRotatorSettings(settings *RotatorSettings) (err error) { return nil } -func getContentLength(rwsc spooledtempfile.ReadWriteSeekCloser) int { +func getContentLength(rwsc spooledtempfile.ReadWriteSeekCloser) int64 { // If the FileName leads to no existing file, it means that the SpooledTempFile // never had the chance to buffer to disk instead of memory, in which case we can // just read the buffer (which should be <= 2MB) and return the length @@ -216,7 +209,7 @@ func getContentLength(rwsc spooledtempfile.ReadWriteSeekCloser) int { rwsc.Seek(0, 0) buf := new(bytes.Buffer) buf.ReadFrom(rwsc) - return buf.Len() + return int64(buf.Len()) } else { // Else, we return the size of the file on disk fileInfo, err := os.Stat(rwsc.FileName()) @@ -224,6 +217,6 @@ func getContentLength(rwsc spooledtempfile.ReadWriteSeekCloser) int { panic(err) } - return int(fileInfo.Size()) + return fileInfo.Size() } } diff --git a/write.go b/write.go index aae04ed..3474e34 100644 --- a/write.go +++ b/write.go @@ -2,6 +2,7 @@ package warc import ( "bufio" + "bytes" "fmt" "io" "strconv" @@ -18,6 +19,12 @@ type Compressor interface { Reset(io.Writer) } +const ( + warcRecordVersionLine = "WARC/1.1\r\n" + warcHeaderEnd = "\r\n" + warcRecordTrailer = "\r\n\r\n" +) + // Writer writes WARC records to WARC files. type Writer struct { Compressor Compressor @@ -59,8 +66,6 @@ type Record struct { func (w *Writer) WriteRecord(r *Record) (recordID string, err error) { defer r.Content.Close() - var written int64 - // Add the mandatories headers if r.Header.Get("WARC-Date") == "" { r.Header.Set("WARC-Date", time.Now().UTC().Format(time.RFC3339Nano)) @@ -75,13 +80,13 @@ func (w *Writer) WriteRecord(r *Record) (recordID string, err error) { r.Header.Set("WARC-Record-ID", "") } - if _, err := io.WriteString(w.BufWriter, "WARC/1.1\r\n"); err != nil { - return recordID, err + var contentLength int64 + if CL := r.Header.Get("Content-Length"); CL != "" { + contentLength, err = strconv.ParseInt(CL, 10, 64) } - - // Write headers - if r.Header.Get("Content-Length") == "" { - r.Header.Set("Content-Length", strconv.Itoa(getContentLength(r.Content))) + if r.Header.Get("Content-Length") == "" || err != nil { + contentLength = getContentLength(r.Content) + r.Header.Set("Content-Length", strconv.FormatInt(contentLength, 10)) } if r.Header.Get("WARC-Block-Digest") == "" { @@ -95,27 +100,21 @@ func (w *Writer) WriteRecord(r *Record) (recordID string, err error) { r.Header.Set("WARC-Block-Digest", digest) } - for key, value := range r.Header { - if _, err := io.WriteString(w.BufWriter, fmt.Sprintf("%s: %s\r\n", key, value)); err != nil { - return recordID, err - } - } - - if _, err := io.WriteString(w.BufWriter, "\r\n"); err != nil { - return recordID, err - } + headerBlock := serializedRecordHeader(r.Header) + w.setCompressorContentSize(int64(len(headerBlock)) + contentLength + int64(len(warcRecordTrailer))) r.Content.Seek(0, 0) - if written, err = io.Copy(w.BufWriter, r.Content); err != nil { + recordReader := io.MultiReader( + bytes.NewReader(headerBlock), + r.Content, + strings.NewReader(warcRecordTrailer), + ) + if _, err = io.Copy(w.BufWriter, recordReader); err != nil { return recordID, err } - if written > 0 { - DataTotal.Add(written) - } - - if _, err := io.WriteString(w.BufWriter, "\r\n\r\n"); err != nil { - return recordID, err + if contentLength > 0 { + DataTotal.Add(int64(contentLength)) } // Flush data @@ -127,6 +126,32 @@ func (w *Writer) WriteRecord(r *Record) (recordID string, err error) { return recordID, nil } +func serializedRecordHeader(header Header) []byte { + var buf bytes.Buffer + buf.WriteString(warcRecordVersionLine) + for key, value := range header { + buf.WriteString(key) + buf.WriteString(": ") + buf.WriteString(value) + buf.WriteString("\r\n") + } + buf.WriteString(warcHeaderEnd) + return buf.Bytes() +} + +// Must be called before any data is written to the compressor. +func (w *Writer) setCompressorContentSize(size int64) { + if w.Compressor == nil { + return + } + + compressor, ok := w.Compressor.(*sizedZstdWriter) + if !ok { + return + } + compressor.SetContentSize(size) +} + // WriteInfoRecord method can be used to write an information record to the WARC file and flush the data func (w *Writer) WriteInfoRecord(payload map[string]string) (recordID string, err error) { // Initialize the record diff --git a/write_zstd_e2e_test.go b/write_zstd_e2e_test.go new file mode 100644 index 0000000..62687b8 --- /dev/null +++ b/write_zstd_e2e_test.go @@ -0,0 +1,289 @@ +package warc + +import ( + "bytes" + "encoding/binary" + "errors" + "io" + "testing" + + "github.com/klauspost/compress/zstd" +) + +func TestZSTDWriterProducesSpecRequiredFrames(t *testing.T) { + var warcBytes bytes.Buffer + + writer, err := NewWriter(&warcBytes, "test.warc.zst", SHA1, CompressionZstd, false, nil) + if err != nil { + t.Fatalf("NewWriter() failed: %v", err) + } + + if _, err := writer.WriteInfoRecord(map[string]string{ + "format": "WARC file version 1.1", + }); err != nil { + t.Fatalf("WriteInfoRecord() failed: %v", err) + } + + payload := makeZSTDTestBytes(1 << 20) + writer.Reset(&warcBytes) + writeZSTDTestRecord(t, writer, "resource", "https://example.com/large", "application/octet-stream", payload) + + writer.Reset(&warcBytes) + writeZSTDTestRecord(t, writer, "metadata", "https://example.com/small", "text/plain", []byte("x")) + + frames := readZSTDFrameInfos(t, warcBytes.Bytes()) + if len(frames) != 3 { + t.Fatalf("frame count = %d, want 3 (warcinfo + large resource + small metadata)", len(frames)) + } + + var largeResourceFrames int + for i, frame := range frames { + if !zstdFrameHasContentSize(frame.compressed) { + t.Fatalf("frame %d is missing Frame_Content_Size", i) + } + if !zstdFrameHasChecksum(frame.compressed) { + t.Fatalf("frame %d is missing Content_Checksum", i) + } + assertZSTDFrameChecksumVerified(t, frame.compressed) + + record := readOnlyWARCRecordFromFrame(t, frame.decoded) + if record.Header.Get("WARC-Type") == "resource" && record.Header.Get("WARC-Target-URI") == "https://example.com/large" { + largeResourceFrames++ + if got := getContentLength(record.Content); got != int64(len(payload)) { + t.Fatalf("resource frame %d content length = %d, want %d", i, got, len(payload)) + } + } + if err := record.Content.Close(); err != nil { + t.Fatalf("close record content: %v", err) + } + } + + if largeResourceFrames != 1 { + t.Fatalf("large resource record used %d ZSTD frames, want 1", largeResourceFrames) + } +} + +func writeZSTDTestRecord(t *testing.T, writer *Writer, recordType, targetURI, contentType string, payload []byte) { + t.Helper() + + record := NewRecord(t.TempDir(), false) + record.Header.Set("WARC-Type", recordType) + record.Header.Set("WARC-Target-URI", targetURI) + record.Header.Set("Content-Type", contentType) + if _, err := record.Content.Write(payload); err != nil { + t.Fatalf("write payload: %v", err) + } + + if _, err := writer.WriteRecord(record); err != nil { + t.Fatalf("WriteRecord(%s): %v", targetURI, err) + } +} + +type zstdFrameInfo struct { + compressed []byte + decoded []byte +} + +func readZSTDFrameInfos(t *testing.T, data []byte) []zstdFrameInfo { + t.Helper() + + dec, err := zstd.NewReader(nil) + if err != nil { + t.Fatalf("zstd.NewReader: %v", err) + } + defer dec.Close() + + var frames []zstdFrameInfo + reader := bytes.NewReader(data) + for { + frame, err := readZSTDFrameBytes(reader) + if errors.Is(err, io.EOF) { + break + } + if err != nil { + t.Fatalf("read ZSTD frame %d: %v", len(frames), err) + } + + decoded, err := dec.DecodeAll(frame, nil) + if err != nil { + t.Fatalf("decode ZSTD frame %d: %v", len(frames), err) + } + frames = append(frames, zstdFrameInfo{ + compressed: frame, + decoded: append([]byte(nil), decoded...), + }) + } + + return frames +} + +func readZSTDFrameBytes(r io.Reader) ([]byte, error) { + var magic [4]byte + n, err := io.ReadFull(r, magic[:]) + if err == io.EOF || err == io.ErrUnexpectedEOF { + if n == 0 { + return nil, io.EOF + } + return nil, err + } + if err != nil { + return nil, err + } + if binary.LittleEndian.Uint32(magic[:]) != 0xfd2fb528 { + return nil, errors.New("invalid ZSTD magic") + } + + frame := append([]byte(nil), magic[:]...) + var fhd [1]byte + if _, err := io.ReadFull(r, fhd[:]); err != nil { + return nil, err + } + frame = append(frame, fhd[0]) + + fcsFlag := (fhd[0] >> 6) & 0x03 + singleSegmentFlag := (fhd[0] >> 5) & 0x01 + contentChecksumFlag := (fhd[0] >> 2) & 0x01 + dictIDFlag := fhd[0] & 0x03 + + headerRestSize := 0 + if singleSegmentFlag == 0 { + headerRestSize++ + } + switch dictIDFlag { + case 1: + headerRestSize++ + case 2: + headerRestSize += 2 + case 3: + headerRestSize += 4 + } + if singleSegmentFlag == 1 && fcsFlag == 0 { + headerRestSize++ + } else { + switch fcsFlag { + case 1: + headerRestSize += 2 + case 2: + headerRestSize += 4 + case 3: + headerRestSize += 8 + } + } + if headerRestSize > 0 { + headerRest := make([]byte, headerRestSize) + if _, err := io.ReadFull(r, headerRest); err != nil { + return nil, err + } + frame = append(frame, headerRest...) + } + + for { + var blockHeader [3]byte + if _, err := io.ReadFull(r, blockHeader[:]); err != nil { + return nil, err + } + frame = append(frame, blockHeader[:]...) + + blockHeaderVal := uint32(blockHeader[0]) | uint32(blockHeader[1])<<8 | uint32(blockHeader[2])<<16 + lastBlock := (blockHeaderVal & 0x01) != 0 + blockType := (blockHeaderVal >> 1) & 0x03 + blockSize := blockHeaderVal >> 3 + if blockType == 3 { + return nil, errors.New("invalid ZSTD block type") + } + + dataSize := blockSize + if blockType == 1 { + dataSize = 1 + } + if dataSize > 0 { + blockData := make([]byte, dataSize) + if _, err := io.ReadFull(r, blockData); err != nil { + return nil, err + } + frame = append(frame, blockData...) + } + if lastBlock { + break + } + } + + if contentChecksumFlag == 1 { + var checksum [4]byte + if _, err := io.ReadFull(r, checksum[:]); err != nil { + return nil, err + } + frame = append(frame, checksum[:]...) + } + + return frame, nil +} + +func zstdFrameHasContentSize(frame []byte) bool { + if len(frame) < 5 || binary.LittleEndian.Uint32(frame[:4]) != 0xfd2fb528 { + return false + } + + fhd := frame[4] + frameContentSizeFlag := (fhd >> 6) & 0x03 + singleSegmentFlag := (fhd >> 5) & 0x01 + return frameContentSizeFlag != 0 || singleSegmentFlag != 0 +} + +func zstdFrameHasChecksum(frame []byte) bool { + if len(frame) < 5 || binary.LittleEndian.Uint32(frame[:4]) != 0xfd2fb528 { + return false + } + + return frame[4]&(1<<2) != 0 +} + +func assertZSTDFrameChecksumVerified(t *testing.T, frame []byte) { + t.Helper() + + corrupt := append([]byte(nil), frame...) + corrupt[len(corrupt)-1] ^= 0xff + + dec, err := zstd.NewReader(nil) + if err != nil { + t.Fatalf("zstd.NewReader: %v", err) + } + defer dec.Close() + + if _, err := dec.DecodeAll(corrupt, nil); err == nil { + t.Fatal("corrupting the ZSTD frame checksum did not fail decoding") + } +} + +func readOnlyWARCRecordFromFrame(t *testing.T, frame []byte) *Record { + t.Helper() + + reader, err := NewReader(bytes.NewReader(frame)) + if err != nil { + t.Fatalf("NewReader(frame): %v", err) + } + defer reader.Close() + + record, err := reader.ReadRecord() + if err != nil { + t.Fatalf("ReadRecord(frame): %v", err) + } + + if extra, err := reader.ReadRecord(); !errors.Is(err, io.EOF) { + if extra != nil { + _ = extra.Content.Close() + } + t.Fatalf("frame decoded to more than one WARC record; second ReadRecord err = %v", err) + } + + return record +} + +func makeZSTDTestBytes(n int) []byte { + b := make([]byte, n) + pattern := []byte("WARC test testtesttesttesttesttesttesttesttesttest") + for i := range b { + b[i] = pattern[i%len(pattern)] + } + return b +} diff --git a/zstd_writer.go b/zstd_writer.go new file mode 100644 index 0000000..d94e40f --- /dev/null +++ b/zstd_writer.go @@ -0,0 +1,42 @@ +package warc + +import ( + "fmt" + "io" + + "github.com/klauspost/compress/zstd" +) + +// sizedZstdWriter wraps a zstd encoder and remembers its current output writer +// so callers can set Frame_Content_Size before writing a record. +type sizedZstdWriter struct { + *zstd.Encoder + output io.Writer +} + +func newSizedZstdWriter(w io.Writer, dictionary []byte) (*sizedZstdWriter, error) { + opts := []zstd.EOption{ + zstd.WithEncoderLevel(zstd.SpeedBetterCompression), + } + if len(dictionary) > 0 { + opts = append(opts, zstd.WithEncoderDict(dictionary)) + } + + enc, err := zstd.NewWriter(w, opts...) + if err != nil { + return nil, fmt.Errorf("creating zstd writer: %w", err) + } + return &sizedZstdWriter{Encoder: enc, output: w}, nil +} + +func (w *sizedZstdWriter) Reset(output io.Writer) { + w.output = output + w.Encoder.Reset(output) +} + +func (w *sizedZstdWriter) SetContentSize(size int64) { + if w.output == nil { + return + } + w.Encoder.ResetContentSize(w.output, size) +}