compress_test.go

  1package middleware
  2
  3import (
  4	"compress/flate"
  5	"compress/gzip"
  6	"fmt"
  7	"io"
  8	"net/http"
  9	"net/http/httptest"
 10	"strings"
 11	"testing"
 12
 13	"github.com/google/brotli/go/cbrotli"
 14	"github.com/klauspost/compress/zstd"
 15)
 16
 17func TestCompressor(t *testing.T) {
 18
 19	compressor := NewCompressor(5, "text/html", "text/css")
 20	t.Logf("encoders: %+v", compressor.encoders)
 21	t.Logf("pooled: %+v", compressor.pooledEncoders)
 22	if len(compressor.encoders) != 1 || len(compressor.pooledEncoders) != 3 {
 23		t.Errorf("gzip, deflate, zstd should be pooled. brotli unpooled.")
 24	}
 25
 26	compressor.SetEncoder("nop", func(w io.Writer, _ int) io.Writer {
 27		return w
 28	})
 29
 30	if len(compressor.encoders) != 2 {
 31		t.Errorf("nop, brotli encoders should be stored in the encoders map")
 32	}
 33	r := NewRouter()
 34
 35	r.Use(compressor.Handler)
 36
 37	r.HandleFunc("GET /gethtml", func(w http.ResponseWriter, r *http.Request) {
 38		w.Header().Set("Content-Type", "text/html")
 39		w.Write([]byte("textstring"))
 40	})
 41
 42	r.HandleFunc("GET /getcss", func(w http.ResponseWriter, r *http.Request) {
 43		w.Header().Set("Content-Type", "text/html")
 44		w.Write([]byte("textstring"))
 45	})
 46
 47	r.HandleFunc("GET /getplain", func(w http.ResponseWriter, r *http.Request) {
 48		w.Header().Set("Content-Type", "text/html")
 49		w.Write([]byte("textstring"))
 50	})
 51
 52	ts := httptest.NewServer(r)
 53	defer ts.Close()
 54
 55	tests := []struct {
 56		name              string
 57		path              string
 58		expectedEncoding  string
 59		acceptedEncodings []string
 60	}{
 61		{
 62			name:              "no expected encodings due to no accepted encodings",
 63			path:              "/gethtml",
 64			acceptedEncodings: nil,
 65			expectedEncoding:  "",
 66		},
 67		{
 68			name:              "no expected encodings due to content type",
 69			path:              "/getplain",
 70			acceptedEncodings: nil,
 71			expectedEncoding:  "",
 72		},
 73		{
 74			name:              "gzip is only encoding",
 75			path:              "/gethtml",
 76			acceptedEncodings: []string{"gzip"},
 77			expectedEncoding:  "gzip",
 78		},
 79		{
 80			name:              "gzip is preferred over deflate",
 81			path:              "/getcss",
 82			acceptedEncodings: []string{"gzip", "deflate"},
 83			expectedEncoding:  "gzip",
 84		},
 85		{
 86			name:              "deflate is used",
 87			path:              "/getcss",
 88			acceptedEncodings: []string{"deflate"},
 89			expectedEncoding:  "deflate",
 90		},
 91		{
 92			name:              "brotli is used",
 93			path:              "/getcss",
 94			acceptedEncodings: []string{"br"},
 95			expectedEncoding:  "br",
 96		},
 97		{
 98			name:              "zstd is used",
 99			path:              "/getcss",
100			acceptedEncodings: []string{"zstd"},
101			expectedEncoding:  "zstd",
102		},
103		{
104
105			name:              "nop is preferred",
106			path:              "/getcss",
107			acceptedEncodings: []string{"nop, gzip, deflate"},
108			expectedEncoding:  "nop",
109		},
110	}
111
112	for _, tc := range tests {
113		tc := tc
114		t.Run(tc.name, func(t *testing.T) {
115			resp, respString := testRequestWithAcceptedEncodings(t, ts, "GET", tc.path, tc.acceptedEncodings...)
116			if respString != "textstring" {
117				t.Errorf("response text doesn't match; expected:%q, got:%q", "textstring", respString)
118			}
119			if got := resp.Header.Get("Content-Encoding"); got != tc.expectedEncoding {
120				t.Errorf("expected encoding %q but got %q", tc.expectedEncoding, got)
121			}
122
123		})
124
125	}
126}
127
128func TestCompressorWildcards(t *testing.T) {
129	tests := []struct {
130		name       string
131		recover    string
132		types      []string
133		typesCount int
134		wcCount    int
135	}{
136		{
137			name:       "defaults",
138			typesCount: 13,
139		},
140		{
141			name:       "no wildcard",
142			types:      []string{"text/plain", "text/html"},
143			typesCount: 2,
144		},
145		{
146			name:    "invalid wildcard #1",
147			types:   []string{"audio/*wav"},
148			recover: "middleware/compress: Unsupported content-type wildcard pattern 'audio/*wav'. Only '/*' supported",
149		},
150		{
151			name:    "invalid wildcard #2",
152			types:   []string{"application*/*"},
153			recover: "middleware/compress: Unsupported content-type wildcard pattern 'application*/*'. Only '/*' supported",
154		},
155		{
156			name:    "valid wildcard",
157			types:   []string{"text/*"},
158			wcCount: 1,
159		},
160		{
161			name:       "mixed",
162			types:      []string{"audio/wav", "text/*"},
163			typesCount: 1,
164			wcCount:    1,
165		},
166	}
167	for _, tt := range tests {
168		t.Run(tt.name, func(t *testing.T) {
169			defer func() {
170				if tt.recover == "" {
171					tt.recover = "<nil>"
172				}
173				if r := recover(); tt.recover != fmt.Sprintf("%v", r) {
174					t.Errorf("Unexpected value recovered: %v", r)
175				}
176			}()
177			compressor := NewCompressor(5, tt.types...)
178			if len(compressor.allowedTypes) != tt.typesCount {
179				t.Errorf("expected %d allowedTypes, got %d", tt.typesCount, len(compressor.allowedTypes))
180			}
181			if len(compressor.allowedWildcards) != tt.wcCount {
182				t.Errorf("expected %d allowedWildcards, got %d", tt.wcCount, len(compressor.allowedWildcards))
183			}
184		})
185	}
186}
187
188func testRequestWithAcceptedEncodings(t *testing.T, ts *httptest.Server, method, path string, encodings ...string) (*http.Response, string) {
189	req, err := http.NewRequest(method, ts.URL+path, nil)
190	if err != nil {
191		t.Fatal(err)
192		return nil, ""
193	}
194	if len(encodings) > 0 {
195		encodingsString := strings.Join(encodings, ",")
196		req.Header.Set("Accept-Encoding", encodingsString)
197	}
198
199	resp, err := http.DefaultClient.Do(req)
200	if err != nil {
201		t.Fatal(err)
202		return nil, ""
203	}
204
205	respBody := decodeResponseBody(t, resp)
206	defer resp.Body.Close()
207
208	return resp, respBody
209}
210
211func decodeResponseBody(t *testing.T, resp *http.Response) string {
212	var reader io.ReadCloser
213	switch resp.Header.Get("Content-Encoding") {
214	case "gzip":
215		var err error
216		reader, err = gzip.NewReader(resp.Body)
217		if err != nil {
218			t.Fatal(err)
219		}
220	case "deflate":
221		reader = flate.NewReader(resp.Body)
222	case "zstd":
223		r, _ := zstd.NewReader(resp.Body)
224		reader = r.IOReadCloser()
225	case "br":
226		reader = cbrotli.NewReader(resp.Body)
227	default:
228		reader = resp.Body
229	}
230	respBody, err := io.ReadAll(reader)
231	if err != nil {
232		t.Fatal(err)
233		return ""
234	}
235	reader.Close()
236
237	return string(respBody)
238}