Files
glchat/internal/httpx/middleware.go
T

171 lines
4.1 KiB
Go
Raw Normal View History

package httpx
import (
"context"
"log/slog"
"net"
"net/http"
"strings"
"time"
)
type Middleware func(http.Handler) http.Handler
func Chain(h http.Handler, mws ...Middleware) http.Handler {
for i := len(mws) - 1; i >= 0; i-- {
h = mws[i](h)
}
return h
}
type statusRecorder struct {
http.ResponseWriter
status int
bytes int
}
func (r *statusRecorder) WriteHeader(status int) {
r.status = status
r.ResponseWriter.WriteHeader(status)
}
func (r *statusRecorder) Write(b []byte) (int, error) {
if r.status == 0 {
r.status = http.StatusOK
}
n, err := r.ResponseWriter.Write(b)
r.bytes += n
return n, err
}
func (r *statusRecorder) Flush() {
if f, ok := r.ResponseWriter.(http.Flusher); ok {
f.Flush()
}
}
func RequestID(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
id := r.Header.Get("X-Request-Id")
if id == "" {
id = newRequestID()
}
w.Header().Set("X-Request-Id", id)
next.ServeHTTP(w, r.WithContext(withRequestID(r.Context(), id)))
})
}
func Logger(logger *slog.Logger) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
start := time.Now()
rec := &statusRecorder{ResponseWriter: w}
next.ServeHTTP(rec, r)
if rec.status == 0 {
rec.status = http.StatusOK
}
logger.LogAttrs(r.Context(), slog.LevelInfo, "http request",
slog.String("request_id", RequestIDFrom(r.Context())),
slog.String("method", r.Method),
slog.String("path", r.URL.Path),
slog.Int("status", rec.status),
slog.Int("bytes", rec.bytes),
slog.String("remote_ip", ClientIP(r, nil)),
slog.Duration("duration", time.Since(start)),
)
})
}
}
func Recoverer(logger *slog.Logger) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer func() {
if rec := recover(); rec != nil {
logger.ErrorContext(r.Context(), "panic recovered",
slog.String("request_id", RequestIDFrom(r.Context())),
slog.String("path", r.URL.Path),
slog.Any("panic", rec),
)
WriteError(w, NewError(CodeInternalError, "internal error"))
}
}()
next.ServeHTTP(w, r)
})
}
}
func SecurityHeaders(filesDomain string) Middleware {
frameAncestors := "'none'"
csp := strings.Join([]string{
"default-src 'self'",
"base-uri 'self'",
"object-src 'none'",
"frame-ancestors " + frameAncestors,
"form-action 'self'",
}, "; ")
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
h := w.Header()
h.Set("Content-Security-Policy", csp)
h.Set("Referrer-Policy", "strict-origin-when-cross-origin")
h.Set("X-Content-Type-Options", "nosniff")
h.Set("X-Frame-Options", "DENY")
h.Set("Cross-Origin-Resource-Policy", "same-site")
if isTLS(r) {
h.Set("Strict-Transport-Security", "max-age=31536000; includeSubDomains")
}
next.ServeHTTP(w, r)
})
}
}
func isTLS(r *http.Request) bool {
return r.TLS != nil || strings.EqualFold(r.Header.Get("X-Forwarded-Proto"), "https")
}
func ClientIP(r *http.Request, trustedProxies []*net.IPNet) string {
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
host = r.RemoteAddr
}
proxyTrusted := false
for _, n := range trustedProxies {
if ip := net.ParseIP(host); ip != nil && n.Contains(ip) {
proxyTrusted = true
break
}
}
if !proxyTrusted {
return host
}
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
parts := strings.Split(xff, ",")
candidate := strings.TrimSpace(parts[len(parts)-1])
if ip := net.ParseIP(candidate); ip != nil {
return ip.String()
}
}
if realIP := strings.TrimSpace(r.Header.Get("X-Real-Ip")); realIP != "" {
if ip := net.ParseIP(realIP); ip != nil {
return ip.String()
}
}
return host
}
type ctxKey int
const requestIDKey ctxKey = iota
func withRequestID(ctx context.Context, id string) context.Context {
return context.WithValue(ctx, requestIDKey, id)
}
func RequestIDFrom(ctx context.Context) string {
if v, ok := ctx.Value(requestIDKey).(string); ok {
return v
}
return ""
}