Source file src/encoding/base64/base64_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 base64
     6  
     7  import (
     8  	"bytes"
     9  	"errors"
    10  	"fmt"
    11  	"io"
    12  	"math"
    13  	"reflect"
    14  	"runtime/debug"
    15  	"strconv"
    16  	"strings"
    17  	"testing"
    18  	"time"
    19  )
    20  
    21  type testpair struct {
    22  	decoded, encoded string
    23  }
    24  
    25  var pairs = []testpair{
    26  	// RFC 3548 examples
    27  	{"\x14\xfb\x9c\x03\xd9\x7e", "FPucA9l+"},
    28  	{"\x14\xfb\x9c\x03\xd9", "FPucA9k="},
    29  	{"\x14\xfb\x9c\x03", "FPucAw=="},
    30  
    31  	// RFC 4648 examples
    32  	{"", ""},
    33  	{"f", "Zg=="},
    34  	{"fo", "Zm8="},
    35  	{"foo", "Zm9v"},
    36  	{"foob", "Zm9vYg=="},
    37  	{"fooba", "Zm9vYmE="},
    38  	{"foobar", "Zm9vYmFy"},
    39  
    40  	// Wikipedia examples
    41  	{"sure.", "c3VyZS4="},
    42  	{"sure", "c3VyZQ=="},
    43  	{"sur", "c3Vy"},
    44  	{"su", "c3U="},
    45  	{"leasure.", "bGVhc3VyZS4="},
    46  	{"easure.", "ZWFzdXJlLg=="},
    47  	{"asure.", "YXN1cmUu"},
    48  	{"sure.", "c3VyZS4="},
    49  }
    50  
    51  // Do nothing to a reference base64 string (leave in standard format)
    52  func stdRef(ref string) string {
    53  	return ref
    54  }
    55  
    56  // Convert a reference string to URL-encoding
    57  func urlRef(ref string) string {
    58  	ref = strings.ReplaceAll(ref, "+", "-")
    59  	ref = strings.ReplaceAll(ref, "/", "_")
    60  	return ref
    61  }
    62  
    63  // Convert a reference string to raw, unpadded format
    64  func rawRef(ref string) string {
    65  	return strings.TrimRight(ref, "=")
    66  }
    67  
    68  // Both URL and unpadding conversions
    69  func rawURLRef(ref string) string {
    70  	return rawRef(urlRef(ref))
    71  }
    72  
    73  const encodeStd = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/"
    74  
    75  // A nonstandard encoding with a funny padding character, for testing
    76  var funnyEncoding = NewEncoding(encodeStd).WithPadding(rune('@'))
    77  
    78  func funnyRef(ref string) string {
    79  	return strings.ReplaceAll(ref, "=", "@")
    80  }
    81  
    82  type encodingTest struct {
    83  	enc  *Encoding           // Encoding to test
    84  	conv func(string) string // Reference string converter
    85  }
    86  
    87  var encodingTests = []encodingTest{
    88  	{StdEncoding, stdRef},
    89  	{URLEncoding, urlRef},
    90  	{RawStdEncoding, rawRef},
    91  	{RawURLEncoding, rawURLRef},
    92  	{funnyEncoding, funnyRef},
    93  	{StdEncoding.Strict(), stdRef},
    94  	{URLEncoding.Strict(), urlRef},
    95  	{RawStdEncoding.Strict(), rawRef},
    96  	{RawURLEncoding.Strict(), rawURLRef},
    97  	{funnyEncoding.Strict(), funnyRef},
    98  }
    99  
   100  var bigtest = testpair{
   101  	"Twas brillig, and the slithy toves",
   102  	"VHdhcyBicmlsbGlnLCBhbmQgdGhlIHNsaXRoeSB0b3Zlcw==",
   103  }
   104  
   105  func testEqual(t *testing.T, msg string, args ...any) bool {
   106  	t.Helper()
   107  	if args[len(args)-2] != args[len(args)-1] {
   108  		t.Errorf(msg, args...)
   109  		return false
   110  	}
   111  	return true
   112  }
   113  
   114  func TestEncode(t *testing.T) {
   115  	for _, p := range pairs {
   116  		for _, tt := range encodingTests {
   117  			got := tt.enc.EncodeToString([]byte(p.decoded))
   118  			testEqual(t, "Encode(%q) = %q, want %q", p.decoded, got, tt.conv(p.encoded))
   119  			dst := tt.enc.AppendEncode([]byte("lead"), []byte(p.decoded))
   120  			testEqual(t, `AppendEncode("lead", %q) = %q, want %q`, p.decoded, string(dst), "lead"+tt.conv(p.encoded))
   121  		}
   122  	}
   123  }
   124  
   125  func TestEncoder(t *testing.T) {
   126  	for _, p := range pairs {
   127  		bb := &strings.Builder{}
   128  		encoder := NewEncoder(StdEncoding, bb)
   129  		encoder.Write([]byte(p.decoded))
   130  		encoder.Close()
   131  		testEqual(t, "Encode(%q) = %q, want %q", p.decoded, bb.String(), p.encoded)
   132  	}
   133  }
   134  
   135  func TestEncoderBuffering(t *testing.T) {
   136  	input := []byte(bigtest.decoded)
   137  	for bs := 1; bs <= 12; bs++ {
   138  		bb := &strings.Builder{}
   139  		encoder := NewEncoder(StdEncoding, bb)
   140  		for pos := 0; pos < len(input); pos += bs {
   141  			end := pos + bs
   142  			if end > len(input) {
   143  				end = len(input)
   144  			}
   145  			n, err := encoder.Write(input[pos:end])
   146  			testEqual(t, "Write(%q) gave error %v, want %v", input[pos:end], err, error(nil))
   147  			testEqual(t, "Write(%q) gave length %v, want %v", input[pos:end], n, end-pos)
   148  		}
   149  		err := encoder.Close()
   150  		testEqual(t, "Close gave error %v, want %v", err, error(nil))
   151  		testEqual(t, "Encoding/%d of %q = %q, want %q", bs, bigtest.decoded, bb.String(), bigtest.encoded)
   152  	}
   153  }
   154  
   155  func TestDecode(t *testing.T) {
   156  	for _, p := range pairs {
   157  		for _, tt := range encodingTests {
   158  			encoded := tt.conv(p.encoded)
   159  			dbuf := make([]byte, tt.enc.DecodedLen(len(encoded)))
   160  			count, err := tt.enc.Decode(dbuf, []byte(encoded))
   161  			testEqual(t, "Decode(%q) = error %v, want %v", encoded, err, error(nil))
   162  			testEqual(t, "Decode(%q) = length %v, want %v", encoded, count, len(p.decoded))
   163  			testEqual(t, "Decode(%q) = %q, want %q", encoded, string(dbuf[0:count]), p.decoded)
   164  
   165  			dbuf, err = tt.enc.DecodeString(encoded)
   166  			testEqual(t, "DecodeString(%q) = error %v, want %v", encoded, err, error(nil))
   167  			testEqual(t, "DecodeString(%q) = %q, want %q", encoded, string(dbuf), p.decoded)
   168  
   169  			dst, err := tt.enc.AppendDecode([]byte("lead"), []byte(encoded))
   170  			testEqual(t, "AppendDecode(%q) = error %v, want %v", p.encoded, err, error(nil))
   171  			testEqual(t, `AppendDecode("lead", %q) = %q, want %q`, p.encoded, string(dst), "lead"+p.decoded)
   172  
   173  			dst2, err := tt.enc.AppendDecode(dst[:0:len(p.decoded)], []byte(encoded))
   174  			testEqual(t, "AppendDecode(%q) = error %v, want %v", p.encoded, err, error(nil))
   175  			testEqual(t, `AppendDecode("", %q) = %q, want %q`, p.encoded, string(dst2), p.decoded)
   176  			if len(dst) > 0 && len(dst2) > 0 && &dst[0] != &dst2[0] {
   177  				t.Errorf("unexpected capacity growth: got %d, want %d", cap(dst2), cap(dst))
   178  			}
   179  		}
   180  	}
   181  }
   182  
   183  func TestDecoder(t *testing.T) {
   184  	for _, p := range pairs {
   185  		decoder := NewDecoder(StdEncoding, strings.NewReader(p.encoded))
   186  		dbuf := make([]byte, StdEncoding.DecodedLen(len(p.encoded)))
   187  		count, err := decoder.Read(dbuf)
   188  		if err != nil && err != io.EOF {
   189  			t.Fatal("Read failed", err)
   190  		}
   191  		testEqual(t, "Read from %q = length %v, want %v", p.encoded, count, len(p.decoded))
   192  		testEqual(t, "Decoding of %q = %q, want %q", p.encoded, string(dbuf[0:count]), p.decoded)
   193  		if err != io.EOF {
   194  			_, err = decoder.Read(dbuf)
   195  		}
   196  		testEqual(t, "Read from %q = %v, want %v", p.encoded, err, io.EOF)
   197  	}
   198  }
   199  
   200  func TestDecoderChunking(t *testing.T) {
   201  	// The decoder must behave identically to decoding the whole input at
   202  	// once, regardless of how the underlying reader chunks the input.
   203  	// See golang.org/issue/31626.
   204  	tests := []struct {
   205  		enc *Encoding
   206  		in  string
   207  	}{
   208  		{StdEncoding, "Rw==bw=="},     // padding inside the stream
   209  		{StdEncoding, "AAAA####"},     // error offset must not reset per chunk
   210  		{StdEncoding, "Rw==x"},        // trailing garbage after a padded group
   211  		{StdEncoding, "Rw===="},       // extra padding after a padded group
   212  		{StdEncoding, "AAAABBBBCCCC"}, // valid input
   213  		{StdEncoding, "AAAABB=="},     // valid input with padding
   214  		{RawStdEncoding, "AAAABB"},    // valid input, no padding
   215  		{RawStdEncoding, "AAAA#B"},    // invalid byte in final fragment
   216  	}
   217  	for _, tt := range tests {
   218  		want, wantErr := tt.enc.DecodeString(tt.in)
   219  		for i := 0; i <= len(tt.in); i++ {
   220  			r := io.MultiReader(strings.NewReader(tt.in[:i]), strings.NewReader(tt.in[i:]))
   221  			got, gotErr := io.ReadAll(NewDecoder(tt.enc, r))
   222  			if !bytes.Equal(got, want) || gotErr != wantErr {
   223  				t.Errorf("Decode(%q) with split at %d = %q, %v; want %q, %v",
   224  					tt.in, i, got, gotErr, want, wantErr)
   225  			}
   226  		}
   227  	}
   228  }
   229  
   230  func TestDecoderBuffering(t *testing.T) {
   231  	for bs := 1; bs <= 12; bs++ {
   232  		decoder := NewDecoder(StdEncoding, strings.NewReader(bigtest.encoded))
   233  		buf := make([]byte, len(bigtest.decoded)+12)
   234  		var total int
   235  		var n int
   236  		var err error
   237  		for total = 0; total < len(bigtest.decoded) && err == nil; {
   238  			n, err = decoder.Read(buf[total : total+bs])
   239  			total += n
   240  		}
   241  		if err != nil && err != io.EOF {
   242  			t.Errorf("Read from %q at pos %d = %d, unexpected error %v", bigtest.encoded, total, n, err)
   243  		}
   244  		testEqual(t, "Decoding/%d of %q = %q, want %q", bs, bigtest.encoded, string(buf[0:total]), bigtest.decoded)
   245  	}
   246  }
   247  
   248  func TestDecodeCorrupt(t *testing.T) {
   249  	testCases := []struct {
   250  		input  string
   251  		offset int // -1 means no corruption.
   252  	}{
   253  		{"", -1},
   254  		{"\n", -1},
   255  		{"AAA=\n", -1},
   256  		{"AAAA\n", -1},
   257  		{"!!!!", 0},
   258  		{"====", 0},
   259  		{"x===", 1},
   260  		{"=AAA", 0},
   261  		{"A=AA", 1},
   262  		{"AA=A", 2},
   263  		{"AA==A", 4},
   264  		{"AAA=AAAA", 4},
   265  		{"AAAAA", 4},
   266  		{"AAAAAA", 4},
   267  		{"A=", 1},
   268  		{"A==", 1},
   269  		{"AA=", 3},
   270  		{"AA==", -1},
   271  		{"AAA=", -1},
   272  		{"AAAA", -1},
   273  		{"AAAAAA=", 7},
   274  		{"YWJjZA=====", 8},
   275  		{"A!\n", 1},
   276  		{"A=\n", 1},
   277  	}
   278  	for _, tc := range testCases {
   279  		dbuf := make([]byte, StdEncoding.DecodedLen(len(tc.input)))
   280  		_, err := StdEncoding.Decode(dbuf, []byte(tc.input))
   281  		if tc.offset == -1 {
   282  			if err != nil {
   283  				t.Error("Decoder wrongly detected corruption in", tc.input)
   284  			}
   285  			continue
   286  		}
   287  		switch err := err.(type) {
   288  		case CorruptInputError:
   289  			testEqual(t, "Corruption in %q at offset %v, want %v", tc.input, int(err), tc.offset)
   290  		default:
   291  			t.Error("Decoder failed to detect corruption in", tc)
   292  		}
   293  	}
   294  }
   295  
   296  func TestDecodeBounds(t *testing.T) {
   297  	var buf [32]byte
   298  	s := StdEncoding.EncodeToString(buf[:])
   299  	defer func() {
   300  		if err := recover(); err != nil {
   301  			t.Fatalf("Decode panicked unexpectedly: %v\n%s", err, debug.Stack())
   302  		}
   303  	}()
   304  	n, err := StdEncoding.Decode(buf[:], []byte(s))
   305  	if n != len(buf) || err != nil {
   306  		t.Fatalf("StdEncoding.Decode = %d, %v, want %d, nil", n, err, len(buf))
   307  	}
   308  }
   309  
   310  func TestEncodedLen(t *testing.T) {
   311  	type test struct {
   312  		enc  *Encoding
   313  		n    int
   314  		want int64
   315  	}
   316  	tests := []test{
   317  		{RawStdEncoding, 0, 0},
   318  		{RawStdEncoding, 1, 2},
   319  		{RawStdEncoding, 2, 3},
   320  		{RawStdEncoding, 3, 4},
   321  		{RawStdEncoding, 7, 10},
   322  		{StdEncoding, 0, 0},
   323  		{StdEncoding, 1, 4},
   324  		{StdEncoding, 2, 4},
   325  		{StdEncoding, 3, 4},
   326  		{StdEncoding, 4, 8},
   327  		{StdEncoding, 7, 12},
   328  	}
   329  	// check overflow
   330  	switch strconv.IntSize {
   331  	case 32:
   332  		tests = append(tests, test{RawStdEncoding, (math.MaxInt-5)/8 + 1, 357913942})
   333  		tests = append(tests, test{RawStdEncoding, math.MaxInt/4*3 + 2, math.MaxInt})
   334  	case 64:
   335  		tests = append(tests, test{RawStdEncoding, (math.MaxInt-5)/8 + 1, 1537228672809129302})
   336  		tests = append(tests, test{RawStdEncoding, math.MaxInt/4*3 + 2, math.MaxInt})
   337  	}
   338  	tests = append(tests, test{StdEncoding, math.MaxInt / 4 * 3, math.MaxInt - 3})
   339  	for _, tt := range tests {
   340  		if got := tt.enc.EncodedLen(tt.n); int64(got) != tt.want {
   341  			t.Errorf("EncodedLen(%d): got %d, want %d", tt.n, got, tt.want)
   342  		}
   343  	}
   344  }
   345  
   346  func TestEncodedLenOverflow(t *testing.T) {
   347  	for _, tt := range []struct {
   348  		enc *Encoding
   349  		n   int
   350  	}{
   351  		{StdEncoding, math.MaxInt/4*3 + 1},
   352  		{RawStdEncoding, math.MaxInt/4*3 + 3},
   353  	} {
   354  		func() {
   355  			defer func() {
   356  				if recover() == nil {
   357  					t.Errorf("EncodedLen(%d) did not panic", tt.n)
   358  				}
   359  			}()
   360  			tt.enc.EncodedLen(tt.n)
   361  		}()
   362  	}
   363  }
   364  
   365  func TestDecodedLen(t *testing.T) {
   366  	type test struct {
   367  		enc  *Encoding
   368  		n    int
   369  		want int64
   370  	}
   371  	tests := []test{
   372  		{RawStdEncoding, 0, 0},
   373  		{RawStdEncoding, 2, 1},
   374  		{RawStdEncoding, 3, 2},
   375  		{RawStdEncoding, 4, 3},
   376  		{RawStdEncoding, 10, 7},
   377  		{StdEncoding, 0, 0},
   378  		{StdEncoding, 4, 3},
   379  		{StdEncoding, 8, 6},
   380  	}
   381  	// check overflow
   382  	switch strconv.IntSize {
   383  	case 32:
   384  		tests = append(tests, test{RawStdEncoding, math.MaxInt/6 + 1, 268435456})
   385  		tests = append(tests, test{RawStdEncoding, math.MaxInt, 1610612735})
   386  	case 64:
   387  		tests = append(tests, test{RawStdEncoding, math.MaxInt/6 + 1, 1152921504606846976})
   388  		tests = append(tests, test{RawStdEncoding, math.MaxInt, 6917529027641081855})
   389  	}
   390  	for _, tt := range tests {
   391  		if got := tt.enc.DecodedLen(tt.n); int64(got) != tt.want {
   392  			t.Errorf("DecodedLen(%d): got %d, want %d", tt.n, got, tt.want)
   393  		}
   394  	}
   395  }
   396  
   397  func TestBig(t *testing.T) {
   398  	n := 3*1000 + 1
   399  	raw := make([]byte, n)
   400  	const alpha = "0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ"
   401  	for i := 0; i < n; i++ {
   402  		raw[i] = alpha[i%len(alpha)]
   403  	}
   404  	encoded := new(bytes.Buffer)
   405  	w := NewEncoder(StdEncoding, encoded)
   406  	nn, err := w.Write(raw)
   407  	if nn != n || err != nil {
   408  		t.Fatalf("Encoder.Write(raw) = %d, %v want %d, nil", nn, err, n)
   409  	}
   410  	err = w.Close()
   411  	if err != nil {
   412  		t.Fatalf("Encoder.Close() = %v want nil", err)
   413  	}
   414  	decoded, err := io.ReadAll(NewDecoder(StdEncoding, encoded))
   415  	if err != nil {
   416  		t.Fatalf("io.ReadAll(NewDecoder(...)): %v", err)
   417  	}
   418  
   419  	if !bytes.Equal(raw, decoded) {
   420  		var i int
   421  		for i = 0; i < len(decoded) && i < len(raw); i++ {
   422  			if decoded[i] != raw[i] {
   423  				break
   424  			}
   425  		}
   426  		t.Errorf("Decode(Encode(%d-byte string)) failed at offset %d", n, i)
   427  	}
   428  }
   429  
   430  func TestNewLineCharacters(t *testing.T) {
   431  	// Each of these should decode to the string "sure", without errors.
   432  	const expected = "sure"
   433  	examples := []string{
   434  		"c3VyZQ==",
   435  		"c3VyZQ==\r",
   436  		"c3VyZQ==\n",
   437  		"c3VyZQ==\r\n",
   438  		"c3VyZ\r\nQ==",
   439  		"c3V\ryZ\nQ==",
   440  		"c3V\nyZ\rQ==",
   441  		"c3VyZ\nQ==",
   442  		"c3VyZQ\n==",
   443  		"c3VyZQ=\n=",
   444  		"c3VyZQ=\r\n\r\n=",
   445  	}
   446  	for _, e := range examples {
   447  		buf, err := StdEncoding.DecodeString(e)
   448  		if err != nil {
   449  			t.Errorf("Decode(%q) failed: %v", e, err)
   450  			continue
   451  		}
   452  		if s := string(buf); s != expected {
   453  			t.Errorf("Decode(%q) = %q, want %q", e, s, expected)
   454  		}
   455  	}
   456  }
   457  
   458  type nextRead struct {
   459  	n   int   // bytes to return
   460  	err error // error to return
   461  }
   462  
   463  // faultInjectReader returns data from source, rate-limited
   464  // and with the errors as written to nextc.
   465  type faultInjectReader struct {
   466  	source string
   467  	nextc  <-chan nextRead
   468  }
   469  
   470  func (r *faultInjectReader) Read(p []byte) (int, error) {
   471  	nr := <-r.nextc
   472  	if len(p) > nr.n {
   473  		p = p[:nr.n]
   474  	}
   475  	n := copy(p, r.source)
   476  	r.source = r.source[n:]
   477  	return n, nr.err
   478  }
   479  
   480  // tests that we don't ignore errors from our underlying reader
   481  func TestDecoderIssue3577(t *testing.T) {
   482  	next := make(chan nextRead, 10)
   483  	wantErr := errors.New("my error")
   484  	next <- nextRead{5, nil}
   485  	next <- nextRead{10, wantErr}
   486  	next <- nextRead{0, wantErr}
   487  	d := NewDecoder(StdEncoding, &faultInjectReader{
   488  		source: "VHdhcyBicmlsbGlnLCBhbmQgdGhlIHNsaXRoeSB0b3Zlcw==", // twas brillig...
   489  		nextc:  next,
   490  	})
   491  	errc := make(chan error, 1)
   492  	go func() {
   493  		_, err := io.ReadAll(d)
   494  		errc <- err
   495  	}()
   496  	select {
   497  	case err := <-errc:
   498  		if err != wantErr {
   499  			t.Errorf("got error %v; want %v", err, wantErr)
   500  		}
   501  	case <-time.After(5 * time.Second):
   502  		t.Errorf("timeout; Decoder blocked without returning an error")
   503  	}
   504  }
   505  
   506  func TestDecoderIssue4779(t *testing.T) {
   507  	encoded := `CP/EAT8AAAEF
   508  AQEBAQEBAAAAAAAAAAMAAQIEBQYHCAkKCwEAAQUBAQEBAQEAAAAAAAAAAQACAwQFBgcICQoLEAAB
   509  BAEDAgQCBQcGCAUDDDMBAAIRAwQhEjEFQVFhEyJxgTIGFJGhsUIjJBVSwWIzNHKC0UMHJZJT8OHx
   510  Y3M1FqKygyZEk1RkRcKjdDYX0lXiZfKzhMPTdePzRieUpIW0lcTU5PSltcXV5fVWZnaGlqa2xtbm
   511  9jdHV2d3h5ent8fX5/cRAAICAQIEBAMEBQYHBwYFNQEAAhEDITESBEFRYXEiEwUygZEUobFCI8FS
   512  0fAzJGLhcoKSQ1MVY3M08SUGFqKygwcmNcLSRJNUoxdkRVU2dGXi8rOEw9N14/NGlKSFtJXE1OT0
   513  pbXF1eX1VmZ2hpamtsbW5vYnN0dXZ3eHl6e3x//aAAwDAQACEQMRAD8A9VSSSSUpJJJJSkkkJ+Tj
   514  1kiy1jCJJDnAcCTykpKkuQ6p/jN6FgmxlNduXawwAzaGH+V6jn/R/wCt71zdn+N/qL3kVYFNYB4N
   515  ji6PDVjWpKp9TSXnvTf8bFNjg3qOEa2n6VlLpj/rT/pf567DpX1i6L1hs9Py67X8mqdtg/rUWbbf
   516  +gkp0kkkklKSSSSUpJJJJT//0PVUkkklKVLq3WMDpGI7KzrNjADtYNXvI/Mqr/Pd/q9W3vaxjnvM
   517  NaCXE9gNSvGPrf8AWS3qmba5jjsJhoB0DAf0NDf6sevf+/lf8Hj0JJATfWT6/dV6oXU1uOLQeKKn
   518  EQP+Hubtfe/+R7Mf/g7f5xcocp++Z11JMCJPgFBxOg7/AOuqDx8I/ikpkXkmSdU8mJIJA/O8EMAy
   519  j+mSARB/17pKVXYWHXjsj7yIex0PadzXMO1zT5KHoNA3HT8ietoGhgjsfA+CSnvvqh/jJtqsrwOv
   520  2b6NGNzXfTYexzJ+nU7/ALkf4P8Awv6P9KvTQQ4AgyDqCF85Pho3CTB7eHwXoH+LT65uZbX9X+o2
   521  bqbPb06551Y4
   522  `
   523  	encodedShort := strings.ReplaceAll(encoded, "\n", "")
   524  
   525  	dec := NewDecoder(StdEncoding, strings.NewReader(encoded))
   526  	res1, err := io.ReadAll(dec)
   527  	if err != nil {
   528  		t.Errorf("ReadAll failed: %v", err)
   529  	}
   530  
   531  	dec = NewDecoder(StdEncoding, strings.NewReader(encodedShort))
   532  	var res2 []byte
   533  	res2, err = io.ReadAll(dec)
   534  	if err != nil {
   535  		t.Errorf("ReadAll failed: %v", err)
   536  	}
   537  
   538  	if !bytes.Equal(res1, res2) {
   539  		t.Error("Decoded results not equal")
   540  	}
   541  }
   542  
   543  func TestDecoderIssue7733(t *testing.T) {
   544  	s, err := StdEncoding.DecodeString("YWJjZA=====")
   545  	want := CorruptInputError(8)
   546  	if !reflect.DeepEqual(want, err) {
   547  		t.Errorf("Error = %v; want CorruptInputError(8)", err)
   548  	}
   549  	if string(s) != "abcd" {
   550  		t.Errorf("DecodeString = %q; want abcd", s)
   551  	}
   552  }
   553  
   554  func TestDecoderIssue15656(t *testing.T) {
   555  	_, err := StdEncoding.Strict().DecodeString("WvLTlMrX9NpYDQlEIFlnDB==")
   556  	want := CorruptInputError(22)
   557  	if !reflect.DeepEqual(want, err) {
   558  		t.Errorf("Error = %v; want CorruptInputError(22)", err)
   559  	}
   560  	_, err = StdEncoding.Strict().DecodeString("WvLTlMrX9NpYDQlEIFlnDA==")
   561  	if err != nil {
   562  		t.Errorf("Error = %v; want nil", err)
   563  	}
   564  	_, err = StdEncoding.DecodeString("WvLTlMrX9NpYDQlEIFlnDB==")
   565  	if err != nil {
   566  		t.Errorf("Error = %v; want nil", err)
   567  	}
   568  }
   569  
   570  func BenchmarkEncodeToString(b *testing.B) {
   571  	data := make([]byte, 8192)
   572  	b.SetBytes(int64(len(data)))
   573  	for i := 0; i < b.N; i++ {
   574  		StdEncoding.EncodeToString(data)
   575  	}
   576  }
   577  
   578  func BenchmarkDecodeString(b *testing.B) {
   579  	sizes := []int{2, 4, 8, 64, 8192}
   580  	benchFunc := func(b *testing.B, benchSize int) {
   581  		data := StdEncoding.EncodeToString(make([]byte, benchSize))
   582  		b.SetBytes(int64(len(data)))
   583  		b.ResetTimer()
   584  		for i := 0; i < b.N; i++ {
   585  			StdEncoding.DecodeString(data)
   586  		}
   587  	}
   588  	for _, size := range sizes {
   589  		b.Run(fmt.Sprintf("%d", size), func(b *testing.B) {
   590  			benchFunc(b, size)
   591  		})
   592  	}
   593  }
   594  
   595  func BenchmarkNewEncoding(b *testing.B) {
   596  	b.SetBytes(int64(len(Encoding{}.decodeMap)))
   597  	for i := 0; i < b.N; i++ {
   598  		e := NewEncoding(encodeStd)
   599  		for _, v := range e.decodeMap {
   600  			_ = v
   601  		}
   602  	}
   603  }
   604  
   605  func TestDecoderRaw(t *testing.T) {
   606  	source := "AAAAAA"
   607  	want := []byte{0, 0, 0, 0}
   608  
   609  	// Direct.
   610  	dec1, err := RawURLEncoding.DecodeString(source)
   611  	if err != nil || !bytes.Equal(dec1, want) {
   612  		t.Errorf("RawURLEncoding.DecodeString(%q) = %x, %v, want %x, nil", source, dec1, err, want)
   613  	}
   614  
   615  	// Through reader. Used to fail.
   616  	r := NewDecoder(RawURLEncoding, bytes.NewReader([]byte(source)))
   617  	dec2, err := io.ReadAll(io.LimitReader(r, 100))
   618  	if err != nil || !bytes.Equal(dec2, want) {
   619  		t.Errorf("reading NewDecoder(RawURLEncoding, %q) = %x, %v, want %x, nil", source, dec2, err, want)
   620  	}
   621  
   622  	// Should work with padding.
   623  	r = NewDecoder(URLEncoding, bytes.NewReader([]byte(source+"==")))
   624  	dec3, err := io.ReadAll(r)
   625  	if err != nil || !bytes.Equal(dec3, want) {
   626  		t.Errorf("reading NewDecoder(URLEncoding, %q) = %x, %v, want %x, nil", source+"==", dec3, err, want)
   627  	}
   628  }
   629  

View as plain text