Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
39 changes: 39 additions & 0 deletions client_test.go
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
package warc

import (
"bytes"
"context"
"crypto/rand"
"crypto/rsa"
Expand Down Expand Up @@ -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)
}
}

Expand Down Expand Up @@ -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)
}
}
}

Expand Down
4 changes: 2 additions & 2 deletions dedupe.go
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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 {
Expand Down
6 changes: 3 additions & 3 deletions dialer.go
Original file line number Diff line number Diff line change
Expand Up @@ -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{
Expand Down Expand Up @@ -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")) {
Expand Down
15 changes: 4 additions & 11 deletions utils.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand Down Expand Up @@ -208,22 +201,22 @@ 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
if rwsc.FileName() == "" {
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())
if err != nil {
panic(err)
}

return int(fileInfo.Size())
return fileInfo.Size()
}
}
73 changes: 49 additions & 24 deletions write.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package warc

import (
"bufio"
"bytes"
"fmt"
"io"
"strconv"
Expand All @@ -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
Expand Down Expand Up @@ -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))
Expand All @@ -75,13 +80,13 @@ func (w *Writer) WriteRecord(r *Record) (recordID string, err error) {
r.Header.Set("WARC-Record-ID", "<urn:uuid:"+recordID+">")
}

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") == "" {
Expand All @@ -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)))
Comment thread
yzqzss marked this conversation as resolved.

r.Content.Seek(0, 0)
if written, err = io.Copy(w.BufWriter, r.Content); err != nil {
recordReader := io.MultiReader(
Comment thread
yzqzss marked this conversation as resolved.
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
Expand All @@ -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 {
Comment thread
yzqzss marked this conversation as resolved.
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)
Comment thread
NGTmeaty marked this conversation as resolved.
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
Expand Down
Loading
Loading