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}