compress.go
1package middleware
2
3import (
4 "bufio"
5 "compress/flate"
6 "compress/gzip"
7 "errors"
8 "fmt"
9 "io"
10 "net"
11 "net/http"
12 "strings"
13 "sync"
14
15 "github.com/google/brotli/go/cbrotli"
16 "github.com/klauspost/compress/zstd"
17)
18
19var defaultCompressibleContentTypes = []string{
20 "text/html",
21 "text/css",
22 "text/plain",
23 "text/javascript",
24 "application/javascript",
25 "application/x-javascript",
26 "application/json",
27 "application/atom+xml",
28 "application/rss+xml",
29 "image/svg+xml",
30 "application/geo+json",
31 "application/vnd.geo+json",
32 "application/gpx+xml",
33}
34
35// Compress is a middleware that compresses response
36// body of a given content types to a data format based
37// on Accept-Encoding request header. It uses a given
38// compression level.
39//
40// NOTE: make sure to set the Content-Type header on your response
41// otherwise this middleware will not compress the response body. For ex, in
42// your handler you should set w.Header().Set("Content-Type", http.DetectContentType(yourBody))
43// or set it manually.
44//
45// Passing a compression level of 5 is sensible value
46func Compress(level int, types ...string) func(next http.Handler) http.Handler {
47 compressor := NewCompressor(level, types...)
48 return compressor.Handler
49}
50
51// Compressor represents a set of encoding configurations.
52type Compressor struct {
53 // The mapping of encoder names to encoder functions.
54 encoders map[string]EncoderFunc
55 // The mapping of pooled encoders to pools.
56 pooledEncoders map[string]*sync.Pool
57 // The set of content types allowed to be compressed.
58 allowedTypes map[string]struct{}
59 allowedWildcards map[string]struct{}
60 // The list of encoders in order of decreasing precedence.
61 encodingPrecedence []string
62 level int // The compression level.
63}
64
65// NewCompressor creates a new Compressor that will handle encoding responses.
66//
67// The level should be one of the ones defined in the flate package.
68// The types are the content types that are allowed to be compressed.
69func NewCompressor(level int, types ...string) *Compressor {
70 // If types are provided, set those as the allowed types. If none are
71 // provided, use the default list.
72 allowedTypes := make(map[string]struct{})
73 allowedWildcards := make(map[string]struct{})
74 if len(types) > 0 {
75 for _, t := range types {
76 if strings.Contains(strings.TrimSuffix(t, "/*"), "*") {
77 panic(fmt.Sprintf("middleware/compress: Unsupported content-type wildcard pattern '%s'. Only '/*' supported", t))
78 }
79 if strings.HasSuffix(t, "/*") {
80 allowedWildcards[strings.TrimSuffix(t, "/*")] = struct{}{}
81 } else {
82 allowedTypes[t] = struct{}{}
83 }
84 }
85 } else {
86 for _, t := range defaultCompressibleContentTypes {
87 allowedTypes[t] = struct{}{}
88 }
89 }
90
91 c := &Compressor{
92 level: level,
93 encoders: make(map[string]EncoderFunc),
94 pooledEncoders: make(map[string]*sync.Pool),
95 allowedTypes: allowedTypes,
96 allowedWildcards: allowedWildcards,
97 }
98
99 // Set the default encoders. The precedence order uses the reverse
100 // ordering that the encoders were added. This means adding new encoders
101 // will move them to the front of the order.
102 //
103 // TODO:
104 // lzma: Opera.
105 // sdch: Chrome, Android. Gzip output + dictionary header.
106
107 // HTTP 1.1 "deflate" (RFC 2616) stands for DEFLATE data (RFC 1951)
108 // wrapped with zlib (RFC 1950). The zlib wrapper uses Adler-32
109 // checksum compared to CRC-32 used in "gzip" and thus is faster.
110 //
111 // But.. some old browsers (MSIE, Safari 5.1) incorrectly expect
112 // raw DEFLATE data only, without the mentioned zlib wrapper.
113 // Because of this major confusion, most modern browsers try it
114 // both ways, first looking for zlib headers.
115 // Quote by Mark Adler: http://stackoverflow.com/a/9186091/385548
116 //
117 // The list of browsers having problems is quite big, see:
118 // http://zoompf.com/blog/2012/02/lose-the-wait-http-compression
119 // https://web.archive.org/web/20120321182910/http://www.vervestudios.co/projects/compression-tests/results
120 //
121 // That's why we prefer gzip over deflate. It's just more reliable
122 // and not significantly slower than deflate.
123 c.SetEncoder("deflate", encoderDeflate)
124
125 // TODO: Exception for old MSIE browsers that can't handle non-HTML?
126 // https://zoompf.com/blog/2012/02/lose-the-wait-http-compression
127 c.SetEncoder("gzip", encoderGzip)
128
129 // Brotli using Google's C library (requires cgo)
130 c.SetEncoder("br", encoderBrotli)
131
132 // Zstd
133 c.SetEncoder("zstd", encoderZstd)
134
135 // NOTE: Not implemented, intentionally:
136 // case "compress": // LZW. Deprecated.
137 // case "bzip2": // Too slow on-the-fly.
138 // case "zopfli": // Too slow on-the-fly.
139 // case "xz": // Too slow on-the-fly.
140 return c
141}
142
143// SetEncoder can be used to set the implementation of a compression algorithm.
144//
145// The encoding should be a standardised identifier. See:
146// https://developer.mozilla.org/en-US/docs/Web/HTTP/Headers/Accept-Encoding
147//
148// For example, add the Brotli algorithm:
149//
150// import brotli_enc "gopkg.in/kothar/brotli-go.v0/enc"
151//
152// compressor := middleware.NewCompressor(5, "text/html")
153// compressor.SetEncoder("br", func(w io.Writer, level int) io.Writer {
154// params := brotli_enc.NewBrotliParams()
155// params.SetQuality(level)
156// return brotli_enc.NewBrotliWriter(params, w)
157// })
158func (c *Compressor) SetEncoder(encoding string, fn EncoderFunc) {
159 encoding = strings.ToLower(encoding)
160 if encoding == "" {
161 panic("the encoding can not be empty")
162 }
163 if fn == nil {
164 panic("attempted to set a nil encoder function")
165 }
166
167 // If we are adding a new encoder that is already registered, we have to
168 // clear that one out first.
169 delete(c.pooledEncoders, encoding)
170 delete(c.encoders, encoding)
171
172 // If the encoder supports Resetting (IoReseterWriter), then it can be pooled.
173 encoder := fn(io.Discard, c.level)
174 if _, ok := encoder.(ioResetterWriter); ok {
175 pool := &sync.Pool{
176 New: func() interface{} {
177 return fn(io.Discard, c.level)
178 },
179 }
180 c.pooledEncoders[encoding] = pool
181 }
182 // If the encoder is not in the pooledEncoders, add it to the normal encoders.
183 if _, ok := c.pooledEncoders[encoding]; !ok {
184 c.encoders[encoding] = fn
185 }
186
187 for i, v := range c.encodingPrecedence {
188 if v == encoding {
189 c.encodingPrecedence = append(c.encodingPrecedence[:i], c.encodingPrecedence[i+1:]...)
190 }
191 }
192
193 c.encodingPrecedence = append([]string{encoding}, c.encodingPrecedence...)
194}
195
196// Handler returns a new middleware that will compress the response based on the
197// current Compressor.
198func (c *Compressor) Handler(next http.Handler) http.Handler {
199 return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
200 encoder, encoding, cleanup := c.selectEncoder(r.Header, w)
201
202 cw := &compressResponseWriter{
203 ResponseWriter: w,
204 w: w,
205 contentTypes: c.allowedTypes,
206 contentWildcards: c.allowedWildcards,
207 encoding: encoding,
208 compressible: false, // determined in post-handler
209 }
210 if encoder != nil {
211 cw.w = encoder
212 }
213 // Re-add the encoder to the pool if applicable.
214 defer cleanup()
215 defer cw.Close()
216
217 next.ServeHTTP(cw, r)
218 })
219}
220
221// selectEncoder returns the encoder, the name of the encoder, and a closer function.
222func (c *Compressor) selectEncoder(h http.Header, w io.Writer) (io.Writer, string, func()) {
223 header := h.Get("Accept-Encoding")
224
225 // Parse the names of all accepted algorithms from the header.
226 accepted := strings.Split(strings.ToLower(header), ",")
227
228 // Find supported encoder by accepted list by precedence
229 for _, name := range c.encodingPrecedence {
230 if matchAcceptEncoding(accepted, name) {
231 if pool, ok := c.pooledEncoders[name]; ok {
232 encoder := pool.Get().(ioResetterWriter)
233 cleanup := func() {
234 pool.Put(encoder)
235 }
236 encoder.Reset(w)
237 return encoder, name, cleanup
238
239 }
240 if fn, ok := c.encoders[name]; ok {
241 return fn(w, c.level), name, func() {}
242 }
243 }
244
245 }
246
247 // No encoder found to match the accepted encoding
248 return nil, "", func() {}
249}
250
251func matchAcceptEncoding(accepted []string, encoding string) bool {
252 for _, v := range accepted {
253 if strings.Contains(v, encoding) {
254 return true
255 }
256 }
257 return false
258}
259
260// An EncoderFunc is a function that wraps the provided io.Writer with a
261// streaming compression algorithm and returns it.
262//
263// In case of failure, the function should return nil.
264type EncoderFunc func(w io.Writer, level int) io.Writer
265
266// Interface for types that allow resetting io.Writers.
267type ioResetterWriter interface {
268 io.Writer
269 Reset(w io.Writer)
270}
271
272type compressResponseWriter struct {
273 http.ResponseWriter
274
275 // The streaming encoder writer to be used if there is one. Otherwise,
276 // this is just the normal writer.
277 w io.Writer
278 contentTypes map[string]struct{}
279 contentWildcards map[string]struct{}
280 encoding string
281 wroteHeader bool
282 compressible bool
283}
284
285func (cw *compressResponseWriter) isCompressible() bool {
286 // Parse the first part of the Content-Type response header.
287 contentType := cw.Header().Get("Content-Type")
288 contentType, _, _ = strings.Cut(contentType, ";")
289
290 // Is the content type compressible?
291 if _, ok := cw.contentTypes[contentType]; ok {
292 return true
293 }
294 if contentType, _, hadSlash := strings.Cut(contentType, "/"); hadSlash {
295 _, ok := cw.contentWildcards[contentType]
296 return ok
297 }
298 return false
299}
300
301func (cw *compressResponseWriter) WriteHeader(code int) {
302 if cw.wroteHeader {
303 cw.ResponseWriter.WriteHeader(code) // Allow multiple calls to propagate.
304 return
305 }
306 cw.wroteHeader = true
307 defer cw.ResponseWriter.WriteHeader(code)
308
309 // Already compressed data?
310 if cw.Header().Get("Content-Encoding") != "" {
311 return
312 }
313
314 if !cw.isCompressible() {
315 cw.compressible = false
316 return
317 }
318
319 if cw.encoding != "" {
320 cw.compressible = true
321 cw.Header().Set("Content-Encoding", cw.encoding)
322 cw.Header().Add("Vary", "Accept-Encoding")
323
324 // The content-length after compression is unknown
325 cw.Header().Del("Content-Length")
326 }
327}
328
329func (cw *compressResponseWriter) Write(p []byte) (int, error) {
330 if !cw.wroteHeader {
331 cw.WriteHeader(http.StatusOK)
332 }
333
334 return cw.writer().Write(p)
335}
336
337func (cw *compressResponseWriter) writer() io.Writer {
338 if cw.compressible {
339 return cw.w
340 }
341 return cw.ResponseWriter
342}
343
344type compressFlusher interface {
345 Flush() error
346}
347
348func (cw *compressResponseWriter) Flush() {
349 if f, ok := cw.writer().(http.Flusher); ok {
350 f.Flush()
351 }
352 // If the underlying writer has a compression flush signature,
353 // call this Flush() method instead
354 if f, ok := cw.writer().(compressFlusher); ok {
355 f.Flush()
356
357 // Also flush the underlying response writer
358 if f, ok := cw.ResponseWriter.(http.Flusher); ok {
359 f.Flush()
360 }
361 }
362}
363
364func (cw *compressResponseWriter) Hijack() (net.Conn, *bufio.ReadWriter, error) {
365 if hj, ok := cw.writer().(http.Hijacker); ok {
366 return hj.Hijack()
367 }
368 return nil, nil, errors.New("middleware: http.Hijacker is unavailable on the writer")
369}
370
371func (cw *compressResponseWriter) Push(target string, opts *http.PushOptions) error {
372 if ps, ok := cw.writer().(http.Pusher); ok {
373 return ps.Push(target, opts)
374 }
375 return errors.New("middleware: http.Pusher is unavailable on the writer")
376}
377
378func (cw *compressResponseWriter) Close() error {
379 if c, ok := cw.writer().(io.WriteCloser); ok {
380 return c.Close()
381 }
382 return errors.New("middleware: io.WriteCloser is unavailable on the writer")
383}
384
385func (cw *compressResponseWriter) Unwrap() http.ResponseWriter {
386 return cw.ResponseWriter
387}
388
389func encoderGzip(w io.Writer, level int) io.Writer {
390 gw, err := gzip.NewWriterLevel(w, level)
391 if err != nil {
392 return nil
393 }
394 return gw
395}
396
397func encoderDeflate(w io.Writer, level int) io.Writer {
398 dw, err := flate.NewWriter(w, level)
399 if err != nil {
400 return nil
401 }
402 return dw
403}
404
405func encoderZstd(w io.Writer, level int) io.Writer {
406 zw, err := zstd.NewWriter(w, zstd.WithEncoderLevel(zstd.EncoderLevelFromZstd(level)))
407 if err != nil {
408 return nil
409 }
410 return zw
411}
412
413func encoderBrotli(w io.Writer, level int) io.Writer {
414 return cbrotli.NewWriter(w, cbrotli.WriterOptions{Quality: level})
415}