summaryrefslogtreecommitdiff
path: root/internal/httpapi/proxy.go
blob: 3cd27970a9e4d94649de5a9f47296a756fe9cbd7 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
package httpapi

import (
	"net"
	"net/http"
	"strings"
)

// TrustProxyHeaders accepts forwarding metadata only from explicitly trusted
// CIDRs. This prevents a direct client from forging HTTPS or audit IP state.
func TrustProxyHeaders(next http.Handler, trustedCIDRs []string, requireHTTPS bool) (http.Handler, error) {
	trusted := make([]*net.IPNet, 0, len(trustedCIDRs))
	for _, value := range trustedCIDRs {
		_, network, err := net.ParseCIDR(value)
		if err != nil {
			return nil, err
		}
		trusted = append(trusted, network)
	}
	return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
		host, _, _ := net.SplitHostPort(r.RemoteAddr)
		remote := net.ParseIP(host)
		trustedPeer := false
		for _, network := range trusted {
			if remote != nil && network.Contains(remote) {
				trustedPeer = true
				break
			}
		}
		if !trustedPeer {
			for _, header := range []string{"Forwarded", "X-Forwarded-For", "X-Forwarded-Host", "X-Forwarded-Port", "X-Forwarded-Proto", "X-Real-IP"} {
				r.Header.Del(header)
			}
		} else if forwarded := strings.TrimSpace(strings.Split(r.Header.Get("X-Forwarded-For"), ",")[0]); net.ParseIP(forwarded) != nil {
			r.RemoteAddr = net.JoinHostPort(forwarded, "0")
		}
		secure := r.TLS != nil || (trustedPeer && strings.EqualFold(strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")), "https"))
		if requireHTTPS && !secure {
			http.Error(w, "HTTPS is required", http.StatusUpgradeRequired)
			return
		}
		next.ServeHTTP(w, r)
	}), nil
}