router.go

 1package middleware
 2
 3import (
 4	"net/http"
 5	"slices"
 6)
 7
 8// Router allows Chi-style middleware chains based on the standard library mux
 9type Router struct {
10	globalChain []func(http.Handler) http.Handler
11	routeChain  []func(http.Handler) http.Handler
12	isSubRouter bool
13	*http.ServeMux
14}
15
16// NewRouter
17func NewRouter(mw ...func(http.Handler) http.Handler) *Router {
18	return &Router{
19		globalChain: mw,
20		ServeMux:    http.NewServeMux(),
21	}
22}
23
24// NewDefaultRouter creates a router with a default set of middleware handlers:
25// - Recover: handle unexpected panics
26// - Logger: canonical structured logger in the request context
27// - CORS: set CORS headers
28// - Compress: compress response bodies (all standard types: HTML, CSS, JS, JSON, GeoJSON)
29func NewDefaultRouter() *Router {
30	return NewRouter(Recover, Logger, CORS, Compress(5))
31}
32
33func (r *Router) Use(mw ...func(http.Handler) http.Handler) {
34	if r.isSubRouter {
35		r.routeChain = append(r.routeChain, mw...)
36	} else {
37		r.globalChain = append(r.globalChain, mw...)
38	}
39}
40
41func (r *Router) Group(fn func(r *Router)) {
42	subRouter := &Router{routeChain: slices.Clone(r.routeChain), isSubRouter: true, ServeMux: r.ServeMux}
43	fn(subRouter)
44}
45
46func (r *Router) HandleFunc(pattern string, h http.HandlerFunc) {
47	r.Handle(pattern, h)
48}
49
50func (r *Router) Handle(pattern string, h http.Handler) {
51	for _, mw := range slices.Backward(r.routeChain) {
52		h = mw(h)
53	}
54	r.ServeMux.Handle(pattern, h)
55}
56
57func (r *Router) ServeHTTP(w http.ResponseWriter, rq *http.Request) {
58	var h http.Handler = r.ServeMux
59
60	for _, mw := range slices.Backward(r.globalChain) {
61		h = mw(h)
62	}
63	h.ServeHTTP(w, rq)
64}