Files
pages/internal/httpx/middleware.go
T
2026-08-15 07:13:00 +00:00

333 lines
9.0 KiB
Go

package httpx
import (
"bufio"
"context"
"crypto/rand"
"encoding/hex"
"errors"
"io"
"log/slog"
"net"
"net/http"
"net/netip"
"runtime/debug"
"strings"
"sync"
"time"
)
// Middleware wraps a handler. Chain applies them so that the first listed runs
// outermost.
type Middleware func(http.Handler) http.Handler
// Chain wraps h with mw, outermost first.
func Chain(h http.Handler, mw ...Middleware) http.Handler {
for i := len(mw) - 1; i >= 0; i-- {
h = mw[i](h)
}
return h
}
type ctxKey int
const (
ctxKeyRequestID ctxKey = iota
ctxKeyLogState
)
// ---------------------------------------------------------------- request id
// RequestIDFrom returns the request id assigned by WithRequestID.
func RequestIDFrom(ctx context.Context) string {
id, _ := ctx.Value(ctxKeyRequestID).(string)
return id
}
// WithRequestID assigns each request an id, echoes it in X-Request-Id, and makes
// it available to handlers and the access log.
//
// An inbound X-Request-Id is adopted only when the peer is a trusted proxy and
// the value is short and printable: it ends up in log records, and an arbitrary
// client-controlled string there is a log-forging primitive.
func WithRequestID(trusted []netip.Prefix) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
id := ""
if in := r.Header.Get("X-Request-Id"); in != "" && sanitaryID(in) {
if addr, ok := peerAddr(r); ok && inAny(addr, trusted) {
id = in
}
}
if id == "" {
id = newID()
}
w.Header().Set("X-Request-Id", id)
ctx := context.WithValue(r.Context(), ctxKeyRequestID, id)
next.ServeHTTP(w, r.WithContext(ctx))
})
}
}
func newID() string {
var b [8]byte
if _, err := rand.Read(b[:]); err != nil {
// crypto/rand cannot fail on any supported platform; if it somehow does,
// an empty id is better than taking down the request.
return ""
}
return hex.EncodeToString(b[:])
}
func sanitaryID(s string) bool {
if len(s) == 0 || len(s) > 64 {
return false
}
for i := 0; i < len(s); i++ {
c := s[i]
if c < 0x21 || c > 0x7e {
return false
}
}
return true
}
// ------------------------------------------------------------- log enrichment
// logState collects attributes that handlers discover mid-request (project,
// deployment, key id) so the single access-log line can carry them.
type logState struct {
mu sync.Mutex
attrs []slog.Attr
}
// LogAttr attaches a key/value pair to this request's access-log line. It is a
// no-op outside the middleware chain, so handlers may call it unconditionally.
//
// Never pass a token, an Authorization header, or any part of either.
func LogAttr(ctx context.Context, key string, value any) {
st, ok := ctx.Value(ctxKeyLogState).(*logState)
if !ok {
return
}
st.mu.Lock()
st.attrs = append(st.attrs, slog.Any(key, value))
st.mu.Unlock()
}
// -------------------------------------------------------------- access logger
// AccessLog emits exactly one record per request.
//
// It logs the method, the URL path, and nothing else from the request: headers
// are never logged (Authorization carries a bearer token) and neither is the raw
// query string. Handlers add their own context with LogAttr.
func AccessLog(log *slog.Logger, trusted []netip.Prefix) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
st := &logState{}
ctx := context.WithValue(r.Context(), ctxKeyLogState, st)
r = r.WithContext(ctx)
rec := &recorder{ResponseWriter: w, status: http.StatusOK}
start := time.Now()
next.ServeHTTP(rec, r)
dur := time.Since(start)
level := slog.LevelInfo
switch {
case rec.status >= 500:
level = slog.LevelError
case rec.status == http.StatusNotFound || rec.status == http.StatusMethodNotAllowed:
// Ordinary outcomes on a public listener: a crawler, a stale link,
// a missing favicon. Logging them as warnings would make warnings
// the bulk of the file and hide the ones that mean something.
level = slog.LevelInfo
case rec.status >= 400:
level = slog.LevelWarn
case r.URL.Path == "/healthz" || r.URL.Path == "/readyz":
level = slog.LevelDebug
}
if !log.Enabled(ctx, level) {
return
}
st.mu.Lock()
extra := st.attrs
st.mu.Unlock()
attrs := make([]slog.Attr, 0, 8+len(extra))
attrs = append(attrs,
slog.String("method", r.Method),
slog.String("path", r.URL.Path),
slog.Int("status", rec.status),
slog.Int64("bytes", rec.written),
slog.Float64("dur_ms", float64(dur.Microseconds())/1000),
)
if ip, ok := ClientIP(r, trusted); ok {
attrs = append(attrs, slog.String("ip", ip.String()))
}
if id := RequestIDFrom(ctx); id != "" {
attrs = append(attrs, slog.String("req_id", id))
}
attrs = append(attrs, extra...)
log.LogAttrs(ctx, level, "request", attrs...)
})
}
}
// ------------------------------------------------------------------ recovery
// Recover turns a handler panic into a 500 instead of tearing down the process,
// logging the stack. The panic value itself never reaches the client.
func Recover(log *slog.Logger) Middleware {
return func(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
defer func() {
v := recover()
if v == nil {
return
}
// ErrAbortHandler is the documented way to abort a response;
// net/http expects to handle it and logs nothing.
if err, ok := v.(error); ok && errors.Is(err, http.ErrAbortHandler) {
panic(v)
}
log.ErrorContext(r.Context(), "handler panic",
"panic", v,
"method", r.Method,
"path", r.URL.Path,
"req_id", RequestIDFrom(r.Context()),
"stack", string(debug.Stack()),
)
if rec, ok := w.(*recorder); ok && rec.wroteHeader {
return // response already begun; nothing safe left to send
}
WriteJSON(w, http.StatusInternalServerError,
newInternalEnvelope(RequestIDFrom(r.Context())))
}()
next.ServeHTTP(w, r)
})
}
}
// ---------------------------------------------------------------- client IP
// ClientIP returns the address to attribute the request to. X-Forwarded-For is
// honoured only when the direct peer is itself a trusted proxy; otherwise any
// client could forge its own source address and defeat per-IP rate limiting.
func ClientIP(r *http.Request, trusted []netip.Prefix) (netip.Addr, bool) {
peer, ok := peerAddr(r)
if !ok {
return netip.Addr{}, false
}
if !inAny(peer, trusted) {
return peer, true
}
// Walk right to left and take the first address that is not itself trusted:
// everything to its right was appended by infrastructure we control, and
// everything to its left may have been forged by the client.
xff := r.Header.Values("X-Forwarded-For")
for i := len(xff) - 1; i >= 0; i-- {
parts := strings.Split(xff[i], ",")
for j := len(parts) - 1; j >= 0; j-- {
addr, err := netip.ParseAddr(strings.TrimSpace(parts[j]))
if err != nil {
continue
}
addr = addr.Unmap()
if !inAny(addr, trusted) {
return addr, true
}
}
}
return peer, true
}
func peerAddr(r *http.Request) (netip.Addr, bool) {
host, _, err := net.SplitHostPort(r.RemoteAddr)
if err != nil {
host = r.RemoteAddr
}
addr, err := netip.ParseAddr(host)
if err != nil {
return netip.Addr{}, false
}
return addr.Unmap(), true
}
func inAny(addr netip.Addr, prefixes []netip.Prefix) bool {
for _, p := range prefixes {
if p.Contains(addr) {
return true
}
}
return false
}
// ------------------------------------------------------------ response record
// recorder observes the status and byte count without altering behaviour.
//
// It implements Unwrap, ReadFrom, Flush and Hijack so that wrapping costs
// nothing: without ReadFrom the static handler would lose the sendfile fast
// path that io.Copy takes when the destination is the raw *http.response.
type recorder struct {
http.ResponseWriter
status int
written int64
wroteHeader bool
}
func (r *recorder) WriteHeader(status int) {
if r.wroteHeader {
return
}
r.status = status
r.wroteHeader = true
r.ResponseWriter.WriteHeader(status)
}
func (r *recorder) Write(b []byte) (int, error) {
if !r.wroteHeader {
r.WriteHeader(http.StatusOK)
}
n, err := r.ResponseWriter.Write(b)
r.written += int64(n)
return n, err
}
func (r *recorder) ReadFrom(src io.Reader) (int64, error) {
if !r.wroteHeader {
r.WriteHeader(http.StatusOK)
}
rf, ok := r.ResponseWriter.(io.ReaderFrom)
if !ok {
n, err := io.Copy(r.ResponseWriter, src)
r.written += n
return n, err
}
n, err := rf.ReadFrom(src)
r.written += n
return n, err
}
func (r *recorder) Unwrap() http.ResponseWriter { return r.ResponseWriter }
func (r *recorder) Flush() {
if f, ok := r.ResponseWriter.(http.Flusher); ok {
if !r.wroteHeader {
r.WriteHeader(http.StatusOK)
}
f.Flush()
}
}
func (r *recorder) Hijack() (net.Conn, *bufio.ReadWriter, error) {
h, ok := r.ResponseWriter.(http.Hijacker)
if !ok {
return nil, nil, errors.New("httpx: ResponseWriter does not support hijacking")
}
return h.Hijack()
}