Source file src/compress/flate/writer_test.go

     1  // Copyright 2012 The Go Authors. All rights reserved.
     2  // Use of this source code is governed by a BSD-style
     3  // license that can be found in the LICENSE file.
     4  
     5  package flate
     6  
     7  import (
     8  	"bytes"
     9  	"fmt"
    10  	"io"
    11  	"math"
    12  	"math/rand"
    13  	"runtime"
    14  	"testing"
    15  )
    16  
    17  func BenchmarkEncode(b *testing.B) {
    18  	doBench(b, func(b *testing.B, buf0 []byte, level, n int) {
    19  		b.StopTimer()
    20  		b.SetBytes(int64(n))
    21  
    22  		buf1 := make([]byte, n)
    23  		for i := 0; i < n; i += len(buf0) {
    24  			if len(buf0) > n-i {
    25  				buf0 = buf0[:n-i]
    26  			}
    27  			copy(buf1[i:], buf0)
    28  		}
    29  		buf0 = nil
    30  		w, err := NewWriter(io.Discard, level)
    31  		if err != nil {
    32  			b.Fatal(err)
    33  		}
    34  		runtime.GC()
    35  		b.StartTimer()
    36  		for i := 0; i < b.N; i++ {
    37  			w.Reset(io.Discard)
    38  			w.Write(buf1)
    39  			w.Close()
    40  		}
    41  	})
    42  }
    43  
    44  func BenchmarkWriterMemUsage(b *testing.B) {
    45  	data := make([]byte, 100000)
    46  
    47  	for level := HuffmanOnly; level <= BestCompression; level++ {
    48  		b.Run(fmt.Sprintf("level=%d", level), func(b *testing.B) {
    49  			var zr *Writer
    50  			var err error
    51  			b.ReportAllocs()
    52  			for b.Loop() {
    53  				zr, err = NewWriter(io.Discard, level)
    54  				if err != nil {
    55  					b.Fatal(err)
    56  				}
    57  				zr.Write(data)
    58  				zr.Close()
    59  			}
    60  		})
    61  	}
    62  }
    63  
    64  // errorWriter is a writer that fails after N writes.
    65  type errorWriter struct {
    66  	N int
    67  }
    68  
    69  func (e *errorWriter) Write(b []byte) (int, error) {
    70  	if e.N <= 0 {
    71  		return 0, io.ErrClosedPipe
    72  	}
    73  	e.N--
    74  	return len(b), nil
    75  }
    76  
    77  // Test if errors from the underlying writer is passed upwards.
    78  func TestWriteError(t *testing.T) {
    79  	t.Parallel()
    80  	buf := new(bytes.Buffer)
    81  	n := 65536
    82  	if !testing.Short() {
    83  		n *= 4
    84  	}
    85  	for i := 0; i < n; i++ {
    86  		fmt.Fprintf(buf, "asdasfasf%d%dfghfgujyut%dyutyu\n", i, i, i)
    87  	}
    88  	in := buf.Bytes()
    89  	// We create our own buffer to control number of writes.
    90  	copyBuffer := make([]byte, 128)
    91  	for l := range 10 {
    92  		for fail := 1; fail <= 256; fail *= 2 {
    93  			// Fail after 'fail' writes
    94  			ew := &errorWriter{N: fail}
    95  			w, err := NewWriter(ew, l)
    96  			if err != nil {
    97  				t.Fatalf("NewWriter: level %d: %v", l, err)
    98  			}
    99  			n, err := io.CopyBuffer(w, struct{ io.Reader }{bytes.NewBuffer(in)}, copyBuffer)
   100  			if err == nil {
   101  				t.Fatalf("Level %d: Expected an error, writer was %#v", l, ew)
   102  			}
   103  			n2, err := w.Write([]byte{1, 2, 2, 3, 4, 5})
   104  			if n2 != 0 {
   105  				t.Fatal("Level", l, "Expected 0 length write, got", n)
   106  			}
   107  			if err == nil {
   108  				t.Fatal("Level", l, "Expected an error")
   109  			}
   110  			err = w.Flush()
   111  			if err == nil {
   112  				t.Fatal("Level", l, "Expected an error on flush")
   113  			}
   114  			err = w.Close()
   115  			if err == nil {
   116  				t.Fatal("Level", l, "Expected an error on close")
   117  			}
   118  
   119  			w.Reset(io.Discard)
   120  			n2, err = w.Write([]byte{1, 2, 3, 4, 5, 6})
   121  			if err != nil {
   122  				t.Fatal("Level", l, "Got unexpected error after reset:", err)
   123  			}
   124  			if n2 == 0 {
   125  				t.Fatal("Level", l, "Got 0 length write, expected > 0")
   126  			}
   127  			if testing.Short() {
   128  				return
   129  			}
   130  		}
   131  	}
   132  }
   133  
   134  // Test if errors from the underlying writer is passed upwards.
   135  func TestWriter_Reset(t *testing.T) {
   136  	buf := new(bytes.Buffer)
   137  	n := 65536
   138  	if !testing.Short() {
   139  		n *= 4
   140  	}
   141  	for i := 0; i < n; i++ {
   142  		fmt.Fprintf(buf, "asdasfasf%d%dfghfgujyut%dyutyu\n", i, i, i)
   143  	}
   144  	in := buf.Bytes()
   145  	for l := range 10 {
   146  		l := l
   147  		if testing.Short() && l > 1 {
   148  			break
   149  		}
   150  		t.Run(fmt.Sprintf("level=%d", l), func(t *testing.T) {
   151  			t.Parallel()
   152  			offset := 1
   153  			if testing.Short() {
   154  				offset = 256
   155  			}
   156  			for ; offset <= 256; offset *= 2 {
   157  				// Fail after 'fail' writes
   158  				w, err := NewWriter(io.Discard, l)
   159  				if err != nil {
   160  					t.Fatalf("NewWriter: level %d: %v", l, err)
   161  				}
   162  				if w.d.fast == nil {
   163  					t.Skip("Not Fast...")
   164  					return
   165  				}
   166  				for i := 0; i < (bufferReset-len(in)-offset-maxMatchOffset)/maxMatchOffset; i++ {
   167  					// skip ahead to where we are close to wrap around...
   168  					w.d.fast.reset()
   169  				}
   170  				w.d.fast.reset()
   171  				_, err = w.Write(in)
   172  				if err != nil {
   173  					t.Fatal(err)
   174  				}
   175  				for range 50 {
   176  					// skip ahead again... This should wrap around...
   177  					w.d.fast.reset()
   178  				}
   179  				w.d.fast.reset()
   180  
   181  				_, err = w.Write(in)
   182  				if err != nil {
   183  					t.Fatal(err)
   184  				}
   185  				for range (math.MaxUint32 - bufferReset) / maxMatchOffset {
   186  					// skip ahead to where we are close to wrap around...
   187  					w.d.fast.reset()
   188  				}
   189  
   190  				_, err = w.Write(in)
   191  				if err != nil {
   192  					t.Fatal(err)
   193  				}
   194  				err = w.Close()
   195  				if err != nil {
   196  					t.Fatal(err)
   197  				}
   198  			}
   199  		})
   200  	}
   201  }
   202  
   203  // Test if two runs produce identical results
   204  // even when writing different sizes to the Writer.
   205  func TestDeterministic(t *testing.T) {
   206  	t.Parallel()
   207  	for i := 0; i <= 9; i++ {
   208  		t.Run(fmt.Sprint("L", i), func(t *testing.T) { testDeterministic(i, t) })
   209  	}
   210  	t.Run("LM2", func(t *testing.T) { testDeterministic(-2, t) })
   211  }
   212  
   213  func testDeterministic(i int, t *testing.T) {
   214  	t.Parallel()
   215  	// Test so much we cross a good number of block boundaries.
   216  	var length = maxStoreBlockSize*30 + 500
   217  	if testing.Short() {
   218  		length /= 10
   219  	}
   220  
   221  	// Create a random, but compressible stream.
   222  	rng := rand.New(rand.NewSource(1))
   223  	t1 := make([]byte, length)
   224  	for i := range t1 {
   225  		t1[i] = byte(rng.Int63() & 7)
   226  	}
   227  
   228  	// Do our first encode.
   229  	var b1 bytes.Buffer
   230  	br := bytes.NewBuffer(t1)
   231  	w, err := NewWriter(&b1, i)
   232  	if err != nil {
   233  		t.Fatal(err)
   234  	}
   235  	// Use a very small prime sized buffer.
   236  	cbuf := make([]byte, 787)
   237  	_, err = io.CopyBuffer(w, struct{ io.Reader }{br}, cbuf)
   238  	if err != nil {
   239  		t.Fatal(err)
   240  	}
   241  	w.Close()
   242  
   243  	// We choose a different buffer size,
   244  	// bigger than a maximum block, and also a prime.
   245  	var b2 bytes.Buffer
   246  	cbuf = make([]byte, 81761)
   247  	br2 := bytes.NewBuffer(t1)
   248  	w2, err := NewWriter(&b2, i)
   249  	if err != nil {
   250  		t.Fatal(err)
   251  	}
   252  	_, err = io.CopyBuffer(w2, struct{ io.Reader }{br2}, cbuf)
   253  	if err != nil {
   254  		t.Fatal(err)
   255  	}
   256  	w2.Close()
   257  
   258  	b1b := b1.Bytes()
   259  	b2b := b2.Bytes()
   260  
   261  	if !bytes.Equal(b1b, b2b) {
   262  		t.Errorf("level %d did not produce deterministic result, result mismatch, len(a) = %d, len(b) = %d", i, len(b1b), len(b2b))
   263  	}
   264  
   265  	// Test using io.WriterTo interface.
   266  	var b3 bytes.Buffer
   267  	br = bytes.NewBuffer(t1)
   268  	w, err = NewWriter(&b3, i)
   269  	if err != nil {
   270  		t.Fatal(err)
   271  	}
   272  	_, err = br.WriteTo(w)
   273  	if err != nil {
   274  		t.Fatal(err)
   275  	}
   276  	w.Close()
   277  
   278  	b3b := b3.Bytes()
   279  	if !bytes.Equal(b1b, b3b) {
   280  		t.Errorf("level %d (io.WriterTo) did not produce deterministic result, result mismatch, len(a) = %d, len(b) = %d", i, len(b1b), len(b3b))
   281  	}
   282  }
   283  
   284  // TestDeflateFast_Reset will test that encoding is consistent
   285  // across a warparound of the table offset.
   286  // See https://github.com/golang/go/issues/34121
   287  func TestDeflateFast_Reset(t *testing.T) {
   288  	buf := new(bytes.Buffer)
   289  	n := 65536
   290  
   291  	for i := 0; i < n; i++ {
   292  		fmt.Fprintf(buf, "asdfasdfasdfasdf%d%dfghfgujyut%dyutyu\n", i, i, i)
   293  	}
   294  	// This is specific to level 1.
   295  	const level = 1
   296  	in := buf.Bytes()
   297  	offset := 1
   298  	if testing.Short() {
   299  		offset = 256
   300  	}
   301  
   302  	// We do an encode with a clean buffer to compare.
   303  	var want bytes.Buffer
   304  	w, err := NewWriter(&want, level)
   305  	if err != nil {
   306  		t.Fatalf("NewWriter: level %d: %v", level, err)
   307  	}
   308  
   309  	// Output written 3 times.
   310  	w.Write(in)
   311  	w.Write(in)
   312  	w.Write(in)
   313  	w.Close()
   314  
   315  	for ; offset <= 256; offset *= 2 {
   316  		w, err := NewWriter(io.Discard, level)
   317  		if err != nil {
   318  			t.Fatalf("NewWriter: level %d: %v", level, err)
   319  		}
   320  
   321  		// Reset until we are right before the wraparound.
   322  		// Each reset adds maxMatchOffset to the offset.
   323  		for i := 0; i < (bufferReset-len(in)-offset-maxMatchOffset)/maxMatchOffset; i++ {
   324  			// skip ahead to where we are close to wrap around...
   325  			w.d.reset(nil)
   326  		}
   327  		var got bytes.Buffer
   328  		w.Reset(&got)
   329  
   330  		// Write 3 times, close.
   331  		for i := 0; i < 3; i++ {
   332  			_, err = w.Write(in)
   333  			if err != nil {
   334  				t.Fatal(err)
   335  			}
   336  		}
   337  		err = w.Close()
   338  		if err != nil {
   339  			t.Fatal(err)
   340  		}
   341  		if !bytes.Equal(got.Bytes(), want.Bytes()) {
   342  			t.Fatalf("output did not match at wraparound, len(want)  = %d, len(got) = %d", want.Len(), got.Len())
   343  		}
   344  	}
   345  }
   346  

View as plain text