Source file src/compress/flate/deflate_test.go

     1  // Copyright 2009 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  	"errors"
    10  	"fmt"
    11  	"internal/testenv"
    12  	"io"
    13  	"math/rand"
    14  	"os"
    15  	"reflect"
    16  	"runtime/debug"
    17  	"strconv"
    18  	"strings"
    19  	"sync"
    20  	"testing"
    21  )
    22  
    23  type deflateTest struct {
    24  	in    []byte
    25  	level int
    26  	out   []byte
    27  }
    28  
    29  type deflateInflateTest struct {
    30  	in []byte
    31  }
    32  
    33  type reverseBitsTest struct {
    34  	in       uint16
    35  	bitCount uint8
    36  	out      uint16
    37  }
    38  
    39  var deflateTests = []*deflateTest{
    40  	0: {[]byte{}, 0, []byte{0x3, 0x0}},
    41  	1: {[]byte{0x11}, BestCompression, []byte{0x12, 0x4, 0xc, 0x0}},
    42  	2: {[]byte{0x11}, BestCompression, []byte{0x12, 0x4, 0xc, 0x0}},
    43  	3: {[]byte{0x11}, BestCompression, []byte{0x12, 0x4, 0xc, 0x0}},
    44  
    45  	4: {[]byte{0x11}, 0, []byte{0x0, 0x1, 0x0, 0xfe, 0xff, 0x11, 0x3, 0x0}},
    46  	5: {[]byte{0x11, 0x12}, 0, []byte{0x0, 0x2, 0x0, 0xfd, 0xff, 0x11, 0x12, 0x3, 0x0}},
    47  	6: {[]byte{0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11}, 0,
    48  		[]byte{0x0, 0x8, 0x0, 0xf7, 0xff, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x3, 0x0},
    49  	},
    50  	7:  {[]byte{}, 1, []byte{0x3, 0x0}},
    51  	8:  {[]byte{0x11}, BestCompression, []byte{0x12, 0x4, 0xc, 0x0}},
    52  	9:  {[]byte{0x11, 0x12}, BestCompression, []byte{0x12, 0x14, 0x2, 0xc, 0x0}},
    53  	10: {[]byte{0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11}, BestCompression, []byte{0x12, 0x84, 0x1, 0xc0, 0x0}},
    54  	11: {[]byte{}, 9, []byte{0x3, 0x0}},
    55  	12: {[]byte{0x11}, 9, []byte{0x12, 0x4, 0xc, 0x0}},
    56  	13: {[]byte{0x11, 0x12}, 9, []byte{0x12, 0x14, 0x2, 0xc, 0x0}},
    57  	14: {[]byte{0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11}, 9, []byte{0x12, 0x84, 0x1, 0xc0, 0x0}},
    58  }
    59  
    60  var deflateInflateTests = []*deflateInflateTest{
    61  	{[]byte{}},
    62  	{[]byte{0x11}},
    63  	{[]byte{0x11, 0x12}},
    64  	{[]byte{0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11, 0x11}},
    65  	{[]byte{0x11, 0x10, 0x13, 0x41, 0x21, 0x21, 0x41, 0x13, 0x87, 0x78, 0x13}},
    66  	{largeDataChunk()},
    67  }
    68  
    69  var reverseBitsTests = []*reverseBitsTest{
    70  	{1, 1, 1},
    71  	{1, 2, 2},
    72  	{1, 3, 4},
    73  	{1, 4, 8},
    74  	{1, 5, 16},
    75  	{17, 5, 17},
    76  	{257, 9, 257},
    77  	{29, 5, 23},
    78  }
    79  
    80  func largeDataChunk() []byte {
    81  	result := make([]byte, 100000)
    82  	for i := range result {
    83  		result[i] = byte(i * i & 0xFF)
    84  	}
    85  	return result
    86  }
    87  
    88  func TestBulkHash4(t *testing.T) {
    89  	for _, x := range deflateTests {
    90  		y := x.out
    91  		if len(y) < minMatchLength {
    92  			continue
    93  		}
    94  		y = append(y, y...)
    95  		for j := 4; j < len(y); j++ {
    96  			y := y[:j]
    97  			dst := make([]uint32, len(y)-minMatchLength+1)
    98  			for i := range dst {
    99  				dst[i] = uint32(i + 100)
   100  			}
   101  			bulkHash4(y, dst)
   102  			for i, got := range dst {
   103  				want := hash4(y[i:])
   104  				if got != want && got == uint32(i)+100 {
   105  					t.Errorf("Len:%d Index:%d, want 0x%08x but not modified", len(y), i, want)
   106  				} else if got != want {
   107  					t.Errorf("Len:%d Index:%d, got 0x%08x want:0x%08x", len(y), i, got, want)
   108  				}
   109  			}
   110  		}
   111  	}
   112  }
   113  
   114  func TestDeflate(t *testing.T) {
   115  	for i, h := range deflateTests {
   116  		var buf bytes.Buffer
   117  		w, err := NewWriter(&buf, h.level)
   118  		if err != nil {
   119  			t.Errorf("NewWriter: %v", err)
   120  			continue
   121  		}
   122  		w.Write(h.in)
   123  		w.Close()
   124  		if !bytes.Equal(buf.Bytes(), h.out) {
   125  			t.Errorf("%d: Deflate(%d, %x) =\n%#v\nwant\n%#v", i, h.level, h.in, buf.Bytes(), h.out)
   126  		}
   127  	}
   128  }
   129  
   130  func TestWriterClose(t *testing.T) {
   131  	b := new(bytes.Buffer)
   132  	zw, err := NewWriter(b, 6)
   133  	if err != nil {
   134  		t.Fatalf("NewWriter: %v", err)
   135  	}
   136  
   137  	if c, err := zw.Write([]byte("Test")); err != nil || c != 4 {
   138  		t.Fatalf("Write to not closed writer: %s, %d", err, c)
   139  	}
   140  
   141  	if err := zw.Close(); err != nil {
   142  		t.Fatalf("Close: %v", err)
   143  	}
   144  
   145  	afterClose := b.Len()
   146  
   147  	if c, err := zw.Write([]byte("Test")); err == nil || c != 0 {
   148  		t.Fatalf("Write to closed writer: %v, %d", err, c)
   149  	}
   150  
   151  	if err := zw.Flush(); err == nil {
   152  		t.Fatalf("Flush to closed writer: %s", err)
   153  	}
   154  
   155  	if err := zw.Close(); err != nil {
   156  		t.Fatalf("Close: %v", err)
   157  	}
   158  
   159  	if afterClose != b.Len() {
   160  		t.Fatalf("Writer wrote data after close. After close: %d. After writes on closed stream: %d", afterClose, b.Len())
   161  	}
   162  }
   163  
   164  // A sparseReader returns a stream consisting of 0s followed by 1<<16 1s.
   165  // This tests missing hash references in a very large input.
   166  type sparseReader struct {
   167  	l   int64
   168  	cur int64
   169  }
   170  
   171  func (r *sparseReader) Read(b []byte) (n int, err error) {
   172  	if r.cur >= r.l {
   173  		return 0, io.EOF
   174  	}
   175  	n = len(b)
   176  	cur := r.cur + int64(n)
   177  	if cur > r.l {
   178  		n -= int(cur - r.l)
   179  		cur = r.l
   180  	}
   181  	for i := range b[0:n] {
   182  		if r.cur+int64(i) >= r.l-1<<16 {
   183  			b[i] = 1
   184  		} else {
   185  			b[i] = 0
   186  		}
   187  	}
   188  	r.cur = cur
   189  	return
   190  }
   191  
   192  func TestVeryLongSparseChunk(t *testing.T) {
   193  	if testing.Short() {
   194  		t.Skip("skipping sparse chunk during short test")
   195  	}
   196  	var buf bytes.Buffer
   197  	w, err := NewWriter(&buf, 1)
   198  	if err != nil {
   199  		t.Errorf("NewWriter: %v", err)
   200  		return
   201  	}
   202  	if _, err = io.Copy(w, &sparseReader{l: 23e8}); err != nil {
   203  		t.Errorf("Compress failed: %v", err)
   204  		return
   205  	}
   206  }
   207  
   208  type syncBuffer struct {
   209  	buf    bytes.Buffer
   210  	mu     sync.RWMutex
   211  	closed bool
   212  	ready  chan bool
   213  }
   214  
   215  func newSyncBuffer() *syncBuffer {
   216  	return &syncBuffer{ready: make(chan bool, 1)}
   217  }
   218  
   219  func (b *syncBuffer) Read(p []byte) (n int, err error) {
   220  	for {
   221  		b.mu.RLock()
   222  		n, err = b.buf.Read(p)
   223  		b.mu.RUnlock()
   224  		if n > 0 || b.closed {
   225  			return
   226  		}
   227  		<-b.ready
   228  	}
   229  }
   230  
   231  func (b *syncBuffer) signal() {
   232  	select {
   233  	case b.ready <- true:
   234  	default:
   235  	}
   236  }
   237  
   238  func (b *syncBuffer) Write(p []byte) (n int, err error) {
   239  	n, err = b.buf.Write(p)
   240  	b.signal()
   241  	return
   242  }
   243  
   244  func (b *syncBuffer) WriteMode() {
   245  	b.mu.Lock()
   246  }
   247  
   248  func (b *syncBuffer) ReadMode() {
   249  	b.mu.Unlock()
   250  	b.signal()
   251  }
   252  
   253  func (b *syncBuffer) Close() error {
   254  	b.closed = true
   255  	b.signal()
   256  	return nil
   257  }
   258  
   259  func testSync(t *testing.T, level int, input []byte, name string) {
   260  	if len(input) == 0 {
   261  		return
   262  	}
   263  
   264  	t.Logf("--testSync %d, %d, %s", level, len(input), name)
   265  	buf := newSyncBuffer()
   266  	buf1 := new(bytes.Buffer)
   267  	buf.WriteMode()
   268  	w, err := NewWriter(io.MultiWriter(buf, buf1), level)
   269  	if err != nil {
   270  		t.Errorf("NewWriter: %v", err)
   271  		return
   272  	}
   273  	r := NewReader(buf)
   274  
   275  	// Write half the input and read back.
   276  	for i := range 2 {
   277  		var lo, hi int
   278  		if i == 0 {
   279  			lo, hi = 0, (len(input)+1)/2
   280  		} else {
   281  			lo, hi = (len(input)+1)/2, len(input)
   282  		}
   283  		t.Logf("#%d: write %d-%d", i, lo, hi)
   284  		if _, err := w.Write(input[lo:hi]); err != nil {
   285  			t.Errorf("testSync: write: %v", err)
   286  			return
   287  		}
   288  		if i == 0 {
   289  			if err := w.Flush(); err != nil {
   290  				t.Errorf("testSync: flush: %v", err)
   291  				return
   292  			}
   293  		} else {
   294  			if err := w.Close(); err != nil {
   295  				t.Errorf("testSync: close: %v", err)
   296  			}
   297  		}
   298  		buf.ReadMode()
   299  		out := make([]byte, hi-lo+1)
   300  		m, err := io.ReadAtLeast(r, out, hi-lo)
   301  		t.Logf("#%d: read %d", i, m)
   302  		if m != hi-lo || err != nil {
   303  			t.Errorf("testSync/%d (%d, %d, %s): read %d: %d, %v (%d left)", i, level, len(input), name, hi-lo, m, err, buf.buf.Len())
   304  			return
   305  		}
   306  		if !bytes.Equal(input[lo:hi], out[:hi-lo]) {
   307  			t.Errorf("testSync/%d: read wrong bytes: %x vs %x", i, input[lo:hi], out[:hi-lo])
   308  			return
   309  		}
   310  		// This test originally checked that after reading
   311  		// the first half of the input, there was nothing left
   312  		// in the read buffer (buf.buf.Len() != 0) but that is
   313  		// not necessarily the case: the write Flush may emit
   314  		// some extra framing bits that are not necessary
   315  		// to process to obtain the first half of the uncompressed
   316  		// data. The test ran correctly most of the time, because
   317  		// the background goroutine had usually read even
   318  		// those extra bits by now, but it's not a useful thing to
   319  		// check.
   320  		buf.WriteMode()
   321  	}
   322  	buf.ReadMode()
   323  	out := make([]byte, 10)
   324  	if n, err := r.Read(out); n > 0 || err != io.EOF {
   325  		t.Errorf("testSync (%d, %d, %s): final Read: %d, %v (hex: %x)", level, len(input), name, n, err, out[0:n])
   326  	}
   327  	if buf.buf.Len() != 0 {
   328  		t.Errorf("testSync (%d, %d, %s): extra data at end", level, len(input), name)
   329  	}
   330  	r.Close()
   331  
   332  	// stream should work for ordinary reader too
   333  	r = NewReader(buf1)
   334  	out, err = io.ReadAll(r)
   335  	if err != nil {
   336  		t.Errorf("testSync: read: %s", err)
   337  		return
   338  	}
   339  	r.Close()
   340  	if !bytes.Equal(input, out) {
   341  		t.Errorf("testSync: decompress(compress(data)) != data: level=%d input=%s", level, name)
   342  	}
   343  }
   344  
   345  func testToFromWithLevelAndLimit(t *testing.T, level int, input []byte, name string, limit int) {
   346  	var buffer bytes.Buffer
   347  	w, err := NewWriter(&buffer, level)
   348  	if err != nil {
   349  		t.Errorf("NewWriter: %v", err)
   350  		return
   351  	}
   352  	w.Write(input)
   353  	w.Close()
   354  	if limit > 0 {
   355  		t.Logf("level: %d - Size:%.2f%%, %d b\n", level, float64(buffer.Len()*100)/float64(limit), buffer.Len())
   356  	}
   357  	if limit > 0 && buffer.Len() > limit {
   358  		t.Errorf("level: %d, len(compress(data)) = %d > limit = %d", level, buffer.Len(), limit)
   359  	}
   360  
   361  	r := NewReader(&buffer)
   362  	out, err := io.ReadAll(r)
   363  	if err != nil {
   364  		t.Errorf("read: %s", err)
   365  		return
   366  	}
   367  	r.Close()
   368  	if !bytes.Equal(input, out) {
   369  		t.Errorf("decompress(compress(data)) != data: level=%d input=%s", level, name)
   370  		return
   371  	}
   372  	testSync(t, level, input, name)
   373  }
   374  
   375  func testToFromWithLimit(t *testing.T, input []byte, name string, limit [11]int) {
   376  	for i := range 10 {
   377  		testToFromWithLevelAndLimit(t, i, input, name, limit[i])
   378  	}
   379  	testToFromWithLevelAndLimit(t, -2, input, name, limit[10])
   380  }
   381  
   382  func TestDeflateInflate(t *testing.T) {
   383  	t.Parallel()
   384  	for i, h := range deflateInflateTests {
   385  		if testing.Short() && len(h.in) > 10000 {
   386  			continue
   387  		}
   388  		testToFromWithLimit(t, h.in, fmt.Sprintf("#%d", i), [11]int{})
   389  	}
   390  }
   391  
   392  func TestReverseBits(t *testing.T) {
   393  	for _, h := range reverseBitsTests {
   394  		if v := reverseBits(h.in, h.bitCount); v != h.out {
   395  			t.Errorf("reverseBits(%v,%v) = %v, want %v",
   396  				h.in, h.bitCount, v, h.out)
   397  		}
   398  	}
   399  }
   400  
   401  type deflateInflateStringTest struct {
   402  	filename string
   403  	label    string
   404  	limit    [11]int // Number 11 is ConstantCompression
   405  }
   406  
   407  var deflateInflateStringTests = []deflateInflateStringTest{
   408  	{
   409  		"../testdata/e.txt",
   410  		"2.718281828...",
   411  		[...]int{100018, 67900, 50960, 51150, 50930, 50790, 50790, 50790, 50790, 50790, 43683 + 100},
   412  	},
   413  	{
   414  		"../../testdata/Isaac.Newton-Opticks.txt",
   415  		"Isaac.Newton-Opticks",
   416  		[...]int{567248, 218338, 201354, 199101, 190627, 182587, 179765, 174982, 173422, 173422, 325240},
   417  	},
   418  }
   419  
   420  func TestDeflateInflateString(t *testing.T) {
   421  	t.Parallel()
   422  	if testing.Short() && testenv.Builder() == "" {
   423  		t.Skip("skipping in short mode")
   424  	}
   425  	for _, test := range deflateInflateStringTests {
   426  		gold, err := os.ReadFile(test.filename)
   427  		if err != nil {
   428  			t.Error(err)
   429  		}
   430  		testToFromWithLimit(t, gold, test.label, test.limit)
   431  
   432  		if testing.Short() {
   433  			break
   434  		}
   435  	}
   436  }
   437  
   438  func TestReaderDict(t *testing.T) {
   439  	const (
   440  		dict = "hello world"
   441  		text = "hello again world"
   442  	)
   443  	var b bytes.Buffer
   444  	w, err := NewWriter(&b, 5)
   445  	if err != nil {
   446  		t.Fatalf("NewWriter: %v", err)
   447  	}
   448  	w.Write([]byte(dict))
   449  	w.Flush()
   450  	b.Reset()
   451  	w.Write([]byte(text))
   452  	w.Close()
   453  
   454  	r := NewReaderDict(&b, []byte(dict))
   455  	data, err := io.ReadAll(r)
   456  	if err != nil {
   457  		t.Fatal(err)
   458  	}
   459  	if string(data) != "hello again world" {
   460  		t.Fatalf("read returned %q want %q", string(data), text)
   461  	}
   462  }
   463  
   464  func TestWriterDict(t *testing.T) {
   465  	const (
   466  		dict = "hello world Lorem ipsum dolor sit amet, consectetur adipiscing elit, sed do eiusmod tempor incididunt ut labore et dolore magna aliqua."
   467  		text = "hello world again Lorem ipsum dolor sit amet"
   468  	)
   469  	// This test is sensitive to algorithm changes that skip
   470  	// data in favour of speed. Higher levels are less prone to this
   471  	// so we test level 4-9.
   472  	for l := 4; l < 9; l++ {
   473  		t.Run(fmt.Sprintf("level=%d", l), func(t *testing.T) {
   474  			var b bytes.Buffer
   475  			w, err := NewWriter(&b, l)
   476  			if err != nil {
   477  				t.Fatalf("NewWriter: %v", err)
   478  			}
   479  			w.Write([]byte(dict))
   480  			w.Flush()
   481  			b.Reset()
   482  			w.Write([]byte(text))
   483  			w.Close()
   484  
   485  			var b1 bytes.Buffer
   486  			w, _ = NewWriterDict(&b1, l, []byte(dict))
   487  			w.Write([]byte(text))
   488  			w.Close()
   489  
   490  			if !bytes.Equal(b1.Bytes(), b.Bytes()) {
   491  				t.Errorf("writer wrote\n%v\n want\n%v", b1.Bytes(), b.Bytes())
   492  			}
   493  		})
   494  	}
   495  }
   496  
   497  // TestNonCompressedBlockDoesntLeakDict checks that the dictionary isn't sent when
   498  // sending a non-compressed block. See https://go.dev/issue/80538
   499  func TestNonCompressedBlockDoesntLeakDict(t *testing.T) {
   500  	data := make([]byte, 763)
   501  	rand.New(rand.NewSource(42)).Read(data)
   502  	dict := []byte("0123456789abcdefghij")
   503  	for l := range BestCompression + 1 {
   504  		t.Run(fmt.Sprintf("level=%d", l), func(t *testing.T) {
   505  			var b bytes.Buffer
   506  			w, err := NewWriterDict(&b, l, dict)
   507  			if err != nil {
   508  				t.Fatalf("NewWriterDict: %v", err)
   509  			}
   510  			if _, err := w.Write(data); err != nil {
   511  				t.Fatalf("Write: %v", err)
   512  			}
   513  			if err := w.Close(); err != nil {
   514  				t.Fatalf("Close: %v", err)
   515  			}
   516  			got, err := io.ReadAll(NewReaderDict(&b, dict))
   517  			if err != nil {
   518  				t.Fatalf("NewReaderDict: %v", err)
   519  			}
   520  			if !bytes.Equal(got, data) {
   521  				t.Errorf("round trip mismatch: got %d bytes, want %d (dictionary emitted: %v)",
   522  					len(got), len(data), bytes.HasPrefix(got, dict))
   523  			}
   524  			if b.Len() != 0 {
   525  				t.Errorf("compressed stream not fully consumed: %d bytes left", b.Len())
   526  			}
   527  		})
   528  	}
   529  }
   530  
   531  // See https://golang.org/issue/2508
   532  func TestRegression2508(t *testing.T) {
   533  	if testing.Short() {
   534  		t.Logf("test disabled with -short")
   535  		return
   536  	}
   537  	w, err := NewWriter(io.Discard, 1)
   538  	if err != nil {
   539  		t.Fatalf("NewWriter: %v", err)
   540  	}
   541  	buf := make([]byte, 1024)
   542  	for range 131072 {
   543  		if _, err := w.Write(buf); err != nil {
   544  			t.Fatalf("writer failed: %v", err)
   545  		}
   546  	}
   547  	w.Close()
   548  }
   549  
   550  func TestWriterReset(t *testing.T) {
   551  	t.Parallel()
   552  	for level := -2; level <= 9; level++ {
   553  		if level == -1 {
   554  			level++
   555  		}
   556  		if testing.Short() && level > 1 {
   557  			break
   558  		}
   559  		w, err := NewWriter(io.Discard, level)
   560  		if err != nil {
   561  			t.Fatalf("NewWriter: %v", err)
   562  		}
   563  		buf := []byte("hello world")
   564  		n := 1024
   565  		if testing.Short() {
   566  			n = 10
   567  		}
   568  		for i := 0; i < n; i++ {
   569  			w.Write(buf)
   570  		}
   571  		w.Reset(io.Discard)
   572  
   573  		wref, err := NewWriter(io.Discard, level)
   574  		if err != nil {
   575  			t.Fatalf("NewWriter: %v", err)
   576  		}
   577  
   578  		// DeepEqual doesn't compare functions.
   579  		w.d.fill, wref.d.fill = nil, nil
   580  		w.d.step, wref.d.step = nil, nil
   581  		w.d.state, wref.d.state = nil, nil
   582  		w.d.fast, wref.d.fast = nil, nil
   583  
   584  		// hashMatch is always overwritten when used.
   585  		if w.d.tokens.n != 0 {
   586  			t.Errorf("level %d Writer not reset after Reset. %d tokens were present", level, w.d.tokens.n)
   587  		}
   588  		// As long as the length is 0, we don't care about the content.
   589  		w.d.tokens = wref.d.tokens
   590  
   591  		// We don't care if there are values in the window, as long as it is at d.index is 0
   592  		w.d.window = wref.d.window
   593  		if !reflect.DeepEqual(w, wref) {
   594  			t.Errorf("level %d Writer not reset after Reset", level)
   595  		}
   596  	}
   597  
   598  	for i := HuffmanOnly; i <= BestCompression; i++ {
   599  		testResetOutput(t, fmt.Sprintf("dict=0/level=%d", i), func(w io.Writer) (*Writer, error) { return NewWriter(w, i) })
   600  	}
   601  	dict := []byte(strings.Repeat("we are the world - how are you?", 3))
   602  	for i := HuffmanOnly; i <= BestCompression; i++ {
   603  		testResetOutput(t, fmt.Sprintf("dict=1/level=%d", i), func(w io.Writer) (*Writer, error) { return NewWriterDict(w, i, dict) })
   604  	}
   605  }
   606  
   607  func testResetOutput(t *testing.T, name string, newWriter func(w io.Writer) (*Writer, error)) {
   608  	t.Run(name, func(t *testing.T) {
   609  		buf := new(bytes.Buffer)
   610  		w, err := newWriter(buf)
   611  		if err != nil {
   612  			t.Fatalf("NewWriter: %v", err)
   613  		}
   614  		b := []byte("hello world - how are you doing?")
   615  		for range 1024 {
   616  			w.Write(b)
   617  		}
   618  		w.Close()
   619  		out1 := buf.Bytes()
   620  
   621  		buf2 := new(bytes.Buffer)
   622  		w.Reset(buf2)
   623  		for range 1024 {
   624  			w.Write(b)
   625  		}
   626  		w.Close()
   627  		out2 := buf2.Bytes()
   628  
   629  		if len(out1) != len(out2) {
   630  			t.Errorf("got %d, expected %d bytes", len(out2), len(out1))
   631  			return
   632  		}
   633  		if !bytes.Equal(out1, out2) {
   634  			mm := 0
   635  			for i, b := range out1[:len(out2)] {
   636  				if b != out2[i] {
   637  					t.Errorf("mismatch index %d: %02x, expected %02x", i, out2[i], b)
   638  				}
   639  				mm++
   640  				if mm == 10 {
   641  					t.Fatal("Stopping")
   642  				}
   643  			}
   644  		}
   645  		t.Logf("got %d bytes", len(out1))
   646  	})
   647  }
   648  
   649  // TestBestSpeed tests that round-tripping through deflate and then inflate
   650  // recovers the original input. The Write sizes are near the thresholds in the
   651  // compressor.encSpeed method (0, 16, 128), as well as near maxStoreBlockSize
   652  // (65535).
   653  func TestBestSpeed(t *testing.T) {
   654  	t.Parallel()
   655  	abc := make([]byte, 128)
   656  	for i := range abc {
   657  		abc[i] = byte(i)
   658  	}
   659  	abcabc := bytes.Repeat(abc, 131072/len(abc))
   660  	var want []byte
   661  
   662  	testCases := [][]int{
   663  		{65536, 0},
   664  		{65536, 1},
   665  		{65536, 1, 256},
   666  		{65536, 1, 65536},
   667  		{65536, 14},
   668  		{65536, 15},
   669  		{65536, 16},
   670  		{65536, 16, 256},
   671  		{65536, 16, 65536},
   672  		{65536, 127},
   673  		{65536, 128},
   674  		{65536, 128, 256},
   675  		{65536, 128, 65536},
   676  		{65536, 129},
   677  		{65536, 65536, 256},
   678  		{65536, 65536, 65536},
   679  	}
   680  
   681  	for i, tc := range testCases {
   682  		if testing.Short() && i >= 6 {
   683  			break
   684  		}
   685  		for _, firstN := range []int{1, 65534, 65535, 65536, 65537, 131072} {
   686  			tc[0] = firstN
   687  		outer:
   688  			for _, flush := range []bool{false, true} {
   689  				buf := new(bytes.Buffer)
   690  				want = want[:0]
   691  
   692  				w, err := NewWriter(buf, BestSpeed)
   693  				if err != nil {
   694  					t.Errorf("i=%d, firstN=%d, flush=%t: NewWriter: %v", i, firstN, flush, err)
   695  					continue
   696  				}
   697  				for _, n := range tc {
   698  					want = append(want, abcabc[:n]...)
   699  					if _, err := w.Write(abcabc[:n]); err != nil {
   700  						t.Errorf("i=%d, firstN=%d, flush=%t: Write: %v", i, firstN, flush, err)
   701  						continue outer
   702  					}
   703  					if !flush {
   704  						continue
   705  					}
   706  					if err := w.Flush(); err != nil {
   707  						t.Errorf("i=%d, firstN=%d, flush=%t: Flush: %v", i, firstN, flush, err)
   708  						continue outer
   709  					}
   710  				}
   711  				if err := w.Close(); err != nil {
   712  					t.Errorf("i=%d, firstN=%d, flush=%t: Close: %v", i, firstN, flush, err)
   713  					continue
   714  				}
   715  
   716  				r := NewReader(buf)
   717  				got, err := io.ReadAll(r)
   718  				if err != nil {
   719  					t.Errorf("i=%d, firstN=%d, flush=%t: ReadAll: %v", i, firstN, flush, err)
   720  					continue
   721  				}
   722  				r.Close()
   723  
   724  				if !bytes.Equal(got, want) {
   725  					t.Errorf("i=%d, firstN=%d, flush=%t: corruption during deflate-then-inflate", i, firstN, flush)
   726  					continue
   727  				}
   728  			}
   729  		}
   730  	}
   731  }
   732  
   733  var errIO = errors.New("IO error")
   734  
   735  // failWriter fails with errIO exactly at the nth call to Write.
   736  type failWriter struct{ n int }
   737  
   738  func (w *failWriter) Write(b []byte) (int, error) {
   739  	w.n--
   740  	if w.n == -1 {
   741  		return 0, errIO
   742  	}
   743  	return len(b), nil
   744  }
   745  
   746  func TestWriterPersistentWriteError(t *testing.T) {
   747  	t.Parallel()
   748  	d, err := os.ReadFile("../../testdata/Isaac.Newton-Opticks.txt")
   749  	if err != nil {
   750  		t.Fatalf("ReadFile: %v", err)
   751  	}
   752  	d = d[:10000] // Keep this test short
   753  
   754  	zw, err := NewWriter(nil, DefaultCompression)
   755  	if err != nil {
   756  		t.Fatalf("NewWriter: %v", err)
   757  	}
   758  
   759  	// Sweep over the threshold at which an error is returned.
   760  	// The variable i makes it such that the ith call to failWriter.Write will
   761  	// return errIO. Since failWriter errors are not persistent, we must ensure
   762  	// that flate.Writer errors are persistent.
   763  	for i := 0; i < 1000; i++ {
   764  		fw := &failWriter{i}
   765  		zw.Reset(fw)
   766  
   767  		_, werr := zw.Write(d)
   768  		cerr := zw.Close()
   769  		ferr := zw.Flush()
   770  		if werr != errIO && werr != nil {
   771  			t.Errorf("test %d, mismatching Write error: got %v, want %v", i, werr, errIO)
   772  		}
   773  		if cerr != errIO && fw.n < 0 {
   774  			t.Errorf("test %d, mismatching Close error: got %v, want %v", i, cerr, errIO)
   775  		}
   776  		if ferr != errIO && fw.n < 0 {
   777  			t.Errorf("test %d, mismatching Flush error: got %v, want %v", i, ferr, errIO)
   778  		}
   779  		if fw.n >= 0 {
   780  			// At this point, the failure threshold was sufficiently high enough
   781  			// that we wrote the whole stream without any errors.
   782  			return
   783  		}
   784  	}
   785  }
   786  func TestWriterPersistentFlushError(t *testing.T) {
   787  	zw, err := NewWriter(&failWriter{0}, DefaultCompression)
   788  	if err != nil {
   789  		t.Fatalf("NewWriter: %v", err)
   790  	}
   791  	flushErr := zw.Flush()
   792  	closeErr := zw.Close()
   793  	_, writeErr := zw.Write([]byte("Test"))
   794  	checkErrors([]error{closeErr, flushErr, writeErr}, errIO, t)
   795  }
   796  
   797  func TestWriterPersistentCloseError(t *testing.T) {
   798  	// If underlying writer return error on closing stream we should persistent this error across all writer calls.
   799  	zw, err := NewWriter(&failWriter{0}, DefaultCompression)
   800  	if err != nil {
   801  		t.Fatalf("NewWriter: %v", err)
   802  	}
   803  	closeErr := zw.Close()
   804  	flushErr := zw.Flush()
   805  	_, writeErr := zw.Write([]byte("Test"))
   806  	checkErrors([]error{closeErr, flushErr, writeErr}, errIO, t)
   807  
   808  	// After closing writer we should persistent "write after close" error across Flush and Write calls, but return nil
   809  	// on next Close calls.
   810  	var b bytes.Buffer
   811  	zw.Reset(&b)
   812  	err = zw.Close()
   813  	if err != nil {
   814  		t.Fatalf("First call to close returned error: %s", err)
   815  	}
   816  	err = zw.Close()
   817  	if err != nil {
   818  		t.Fatalf("Second call to close returned error: %s", err)
   819  	}
   820  
   821  	flushErr = zw.Flush()
   822  	_, writeErr = zw.Write([]byte("Test"))
   823  	checkErrors([]error{flushErr, writeErr}, errWriterClosed, t)
   824  }
   825  
   826  func checkErrors(got []error, want error, t *testing.T) {
   827  	t.Helper()
   828  	for _, err := range got {
   829  		if err != want {
   830  			t.Errorf("Error doesn't match\nWant: %s\nGot: %s", want, got)
   831  		}
   832  	}
   833  }
   834  
   835  func TestBestSpeedMatch(t *testing.T) {
   836  	t.Parallel()
   837  	cases := []struct {
   838  		previous []byte
   839  		current  []byte
   840  		t        int
   841  		s        int
   842  		want     int32
   843  	}{{
   844  		previous: []byte{0, 0, 0, 1, 2},
   845  		current:  []byte{3, 4, 5, 0, 1, 2, 3, 4, 5},
   846  		t:        -3,
   847  		s:        3,
   848  		want:     6,
   849  	}, {
   850  		previous: []byte{0, 0, 0, 1, 2},
   851  		current:  []byte{2, 4, 5, 0, 1, 2, 3, 4, 5},
   852  		t:        -3,
   853  		s:        3,
   854  		want:     3,
   855  	}, {
   856  		previous: []byte{0, 0, 0, 1, 1},
   857  		current:  []byte{3, 4, 5, 0, 1, 2, 3, 4, 5},
   858  		t:        -3,
   859  		s:        3,
   860  		want:     2,
   861  	}, {
   862  		previous: []byte{0, 0, 0, 1, 2},
   863  		current:  []byte{2, 2, 2, 2, 1, 2, 3, 4, 5},
   864  		t:        -1,
   865  		s:        0,
   866  		want:     4,
   867  	}, {
   868  		previous: []byte{0, 0, 0, 1, 2, 3, 4, 5, 2, 2},
   869  		current:  []byte{2, 2, 2, 2, 1, 2, 3, 4, 5},
   870  		t:        -7,
   871  		s:        4,
   872  		want:     5,
   873  	}, {
   874  		previous: []byte{9, 9, 9, 9, 9},
   875  		current:  []byte{2, 2, 2, 2, 1, 2, 3, 4, 5},
   876  		t:        -1,
   877  		s:        0,
   878  		want:     0,
   879  	}, {
   880  		previous: []byte{9, 9, 9, 9, 9},
   881  		current:  []byte{9, 2, 2, 2, 1, 2, 3, 4, 5},
   882  		t:        0,
   883  		s:        1,
   884  		want:     0,
   885  	}, {
   886  		previous: []byte{},
   887  		current:  []byte{2, 2, 2, 2, 1, 2, 3, 4, 5},
   888  		t:        0,
   889  		s:        1,
   890  		want:     3,
   891  	}, {
   892  		previous: []byte{3, 4, 5},
   893  		current:  []byte{3, 4, 5},
   894  		t:        -3,
   895  		s:        0,
   896  		want:     3,
   897  	}, {
   898  		previous: make([]byte, 1000),
   899  		current:  make([]byte, 1000),
   900  		t:        -1000,
   901  		s:        0,
   902  		want:     maxMatchLength - 4,
   903  	}, {
   904  		previous: make([]byte, 200),
   905  		current:  make([]byte, 500),
   906  		t:        -200,
   907  		s:        0,
   908  		want:     maxMatchLength - 4,
   909  	}, {
   910  		previous: make([]byte, 200),
   911  		current:  make([]byte, 500),
   912  		t:        0,
   913  		s:        1,
   914  		want:     maxMatchLength - 4,
   915  	}, {
   916  		previous: make([]byte, maxMatchLength-4),
   917  		current:  make([]byte, 500),
   918  		t:        -(maxMatchLength - 4),
   919  		s:        0,
   920  		want:     maxMatchLength - 4,
   921  	}, {
   922  		previous: make([]byte, 200),
   923  		current:  make([]byte, 500),
   924  		t:        -200,
   925  		s:        400,
   926  		want:     100,
   927  	}, {
   928  		previous: make([]byte, 10),
   929  		current:  make([]byte, 500),
   930  		t:        200,
   931  		s:        400,
   932  		want:     100,
   933  	}}
   934  	for i, c := range cases {
   935  		t.Run(strconv.Itoa(i), func(t *testing.T) {
   936  			var e fastGen
   937  			e.addBlock(c.previous)
   938  			e.addBlock(c.current)
   939  			got := e.matchLenLimited(c.s+len(c.previous), c.t+len(c.previous), e.hist)
   940  			if got != c.want {
   941  				t.Errorf("Test %d: match length, want %d, got %d", i, c.want, got)
   942  			}
   943  		})
   944  	}
   945  }
   946  
   947  func TestBestSpeedMaxMatchOffset(t *testing.T) {
   948  	t.Parallel()
   949  	const abc, xyz = "abcdefgh", "stuvwxyz"
   950  	const inputMargin = 16 - 1
   951  	for _, matchBefore := range []bool{false, true} {
   952  		for _, extra := range []int{0, inputMargin - 1, inputMargin, inputMargin + 1, 2 * inputMargin} {
   953  			for offsetAdj := -5; offsetAdj <= +5; offsetAdj++ {
   954  				report := func(desc string, err error) {
   955  					t.Errorf("matchBefore=%t, extra=%d, offsetAdj=%d: %s%v",
   956  						matchBefore, extra, offsetAdj, desc, err)
   957  				}
   958  
   959  				offset := maxMatchOffset + offsetAdj
   960  
   961  				// Make src to be a []byte of the form
   962  				//	"%s%s%s%s%s" % (abc, zeros0, xyzMaybe, abc, zeros1)
   963  				// where:
   964  				//	zeros0 is approximately maxMatchOffset zeros.
   965  				//	xyzMaybe is either xyz or the empty string.
   966  				//	zeros1 is between 0 and 30 zeros.
   967  				// The difference between the two abc's will be offset, which
   968  				// is maxMatchOffset plus or minus a small adjustment.
   969  				src := make([]byte, offset+len(abc)+extra)
   970  				copy(src, abc)
   971  				if !matchBefore {
   972  					copy(src[offset-len(xyz):], xyz)
   973  				}
   974  				copy(src[offset:], abc)
   975  
   976  				buf := new(bytes.Buffer)
   977  				w, err := NewWriter(buf, BestSpeed)
   978  				if err != nil {
   979  					report("NewWriter: ", err)
   980  					continue
   981  				}
   982  				if _, err := w.Write(src); err != nil {
   983  					report("Write: ", err)
   984  					continue
   985  				}
   986  				if err := w.Close(); err != nil {
   987  					report("Writer.Close: ", err)
   988  					continue
   989  				}
   990  
   991  				r := NewReader(buf)
   992  				dst, err := io.ReadAll(r)
   993  				r.Close()
   994  				if err != nil {
   995  					report("ReadAll: ", err)
   996  					continue
   997  				}
   998  
   999  				if !bytes.Equal(dst, src) {
  1000  					report("", fmt.Errorf("bytes differ after round-tripping"))
  1001  					continue
  1002  				}
  1003  			}
  1004  		}
  1005  	}
  1006  }
  1007  
  1008  type canGetFastGen interface {
  1009  	getFastGen() *fastGen
  1010  }
  1011  
  1012  func TestBestSpeedShiftOffsets(t *testing.T) {
  1013  	// Test if shiftoffsets properly preserves matches and resets out-of-range matches
  1014  	// seen in https://github.com/golang/go/issues/4142
  1015  
  1016  	for level := 1; level <= 6; level++ {
  1017  		t.Run(fmt.Sprintf("level=%d", level), func(t *testing.T) {
  1018  			enc := newFastEnc(level)
  1019  
  1020  			// testData may not generate internal matches.
  1021  			testData := make([]byte, 100)
  1022  			rng := rand.New(rand.NewSource(0))
  1023  			for i := range testData {
  1024  				testData[i] = byte(rng.Uint32())
  1025  			}
  1026  			valOrLen := func(val uint16) uint16 {
  1027  				if val == 0 {
  1028  					return uint16(len(testData))
  1029  				}
  1030  				return val
  1031  			}
  1032  			// Encode the testdata with clean state.
  1033  			// Second part should pick up matches from the first block.
  1034  			var firstTokens, secondTokens tokens
  1035  			enc.encode(&firstTokens, testData)
  1036  			enc.encode(&secondTokens, testData)
  1037  			wantFirstTokens := valOrLen(firstTokens.n)
  1038  			wantSecondTokens := valOrLen(secondTokens.n)
  1039  
  1040  			if wantFirstTokens <= wantSecondTokens {
  1041  				t.Fatalf("test needs matches between inputs to be generated, %d == %d", wantFirstTokens, wantSecondTokens)
  1042  			}
  1043  			// Forward the current indicator to before wraparound.
  1044  			fg := enc.(canGetFastGen).getFastGen()
  1045  			fg.hist = nil
  1046  			fg.cur = bufferReset - int32(len(testData))
  1047  
  1048  			// Part 1 before wrap, should match clean state.
  1049  			var gotTokens tokens
  1050  			enc.encode(&gotTokens, testData)
  1051  			got := valOrLen(gotTokens.n)
  1052  			if wantFirstTokens != got {
  1053  				t.Errorf("got %d, want %d tokens", got, wantFirstTokens)
  1054  			}
  1055  
  1056  			// Verify we are about to wrap.
  1057  			gotCur := int(fg.cur) + len(fg.hist)
  1058  			if gotCur != bufferReset {
  1059  				t.Errorf("got %d, want e.cur to be at bufferReset (%d)", gotCur, bufferReset)
  1060  			}
  1061  
  1062  			// Part 2 should match clean state as well even if wrapped.
  1063  			gotTokens.Reset()
  1064  			enc.encode(&gotTokens, testData)
  1065  			got = valOrLen(gotTokens.n)
  1066  			if wantSecondTokens != got {
  1067  				t.Errorf("got %d, want %d token", got, wantSecondTokens)
  1068  			}
  1069  
  1070  			// Verify that we wrapped.
  1071  			if fg.cur >= bufferReset {
  1072  				t.Errorf("want e.cur to be < bufferReset (%d), got %d", bufferReset, fg.cur)
  1073  			}
  1074  
  1075  			// Forward the current buffer, leaving the matches at the bottom.
  1076  			fg.cur = bufferReset
  1077  			fg.hist = nil
  1078  
  1079  			// Ensure that no matches were picked up.
  1080  			gotTokens.Reset()
  1081  			enc.encode(&gotTokens, testData)
  1082  			got = valOrLen(gotTokens.n)
  1083  			if wantFirstTokens != got {
  1084  				t.Errorf("got %d, want %d tokens", got, wantFirstTokens)
  1085  			}
  1086  		})
  1087  	}
  1088  }
  1089  
  1090  func TestMaxStackSize(t *testing.T) {
  1091  	// This test must not run in parallel with other tests as debug.SetMaxStack
  1092  	// affects all goroutines.
  1093  	n := debug.SetMaxStack(1 << 16)
  1094  	defer debug.SetMaxStack(n)
  1095  
  1096  	var wg sync.WaitGroup
  1097  	defer wg.Wait()
  1098  
  1099  	b := make([]byte, 1<<20)
  1100  	for level := HuffmanOnly; level <= BestCompression; level++ {
  1101  		// Run in separate goroutine to increase probability of stack regrowth.
  1102  		wg.Add(1)
  1103  		go func(level int) {
  1104  			defer wg.Done()
  1105  			zw, err := NewWriter(io.Discard, level)
  1106  			if err != nil {
  1107  				t.Errorf("level %d, NewWriter() = %v, want nil", level, err)
  1108  			}
  1109  			if n, err := zw.Write(b); n != len(b) || err != nil {
  1110  				t.Errorf("level %d, Write() = (%d, %v), want (%d, nil)", level, n, err, len(b))
  1111  			}
  1112  			if err := zw.Close(); err != nil {
  1113  				t.Errorf("level %d, Close() = %v, want nil", level, err)
  1114  			}
  1115  			zw.Reset(io.Discard)
  1116  		}(level)
  1117  	}
  1118  }
  1119  

View as plain text