diff --git a/pkg/probod/probod.go b/pkg/probod/probod.go index 4794c6d77..dc2de0b62 100644 --- a/pkg/probod/probod.go +++ b/pkg/probod/probod.go @@ -64,6 +64,7 @@ import ( "go.probo.inc/probo/pkg/probo" "go.probo.inc/probo/pkg/securecookie" "go.probo.inc/probo/pkg/server" + "go.probo.inc/probo/pkg/server/trustedproxy" "go.probo.inc/probo/pkg/slack" "go.probo.inc/probo/pkg/trust" "go.probo.inc/probo/pkg/webhook" @@ -768,6 +769,9 @@ func (impl *Implm) runApiServer( ctx, span := tracer.Start(ctx, "probod.runApiServer") defer span.End() + trustedProxies := parseIPs(impl.cfg.Api.ProxyProtocol.TrustedProxies) + handler = trustedproxy.NewMiddleware(trustedProxies)(handler) + apiServer := httpserver.NewServer( impl.cfg.Api.Addr, handler, diff --git a/pkg/server/api/clientip/clientip.go b/pkg/server/api/clientip/clientip.go index 8c79b9b5f..dd465ebe2 100644 --- a/pkg/server/api/clientip/clientip.go +++ b/pkg/server/api/clientip/clientip.go @@ -15,35 +15,11 @@ package clientip import ( - "context" "net" "net/http" "strings" ) -type ctxKey struct{} - -// NewMiddleware returns an HTTP middleware that extracts the client IP -// from standard proxy headers and stores it in the request context. -func NewMiddleware() func(http.Handler) http.Handler { - return func(next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - ip := Extract(r) - ctx := context.WithValue(r.Context(), ctxKey{}, ip) - next.ServeHTTP(w, r.WithContext(ctx)) - }) - } -} - -// FromContext returns the client IP stored by the middleware, or an -// empty string if the middleware has not run. -func FromContext(ctx context.Context) string { - if ip, ok := ctx.Value(ctxKey{}).(string); ok { - return ip - } - return "" -} - // Extract resolves the client IP address from standard proxy headers // in priority order: RFC 7239 Forwarded, then X-Forwarded-For, then // the connection's remote address. diff --git a/pkg/server/api/clientip/clientip_test.go b/pkg/server/api/clientip/clientip_test.go new file mode 100644 index 000000000..bcd116b1c --- /dev/null +++ b/pkg/server/api/clientip/clientip_test.go @@ -0,0 +1,125 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package clientip_test + +import ( + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "go.probo.inc/probo/pkg/server/api/clientip" +) + +func TestExtract(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + remoteAddr string + headers map[string]string + want string + }{ + { + name: "remote addr only", + remoteAddr: "192.168.1.1:12345", + want: "192.168.1.1", + }, + { + name: "remote addr without port", + remoteAddr: "192.168.1.1", + want: "192.168.1.1", + }, + { + name: "x-forwarded-for single ip", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{"X-Forwarded-For": "203.0.113.50"}, + want: "203.0.113.50", + }, + { + name: "x-forwarded-for chain", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{"X-Forwarded-For": "203.0.113.50, 70.41.3.18, 150.172.238.178"}, + want: "203.0.113.50", + }, + { + name: "x-forwarded-for with port", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{"X-Forwarded-For": "203.0.113.50:8080"}, + want: "203.0.113.50", + }, + { + name: "forwarded header simple", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{"Forwarded": "for=198.51.100.17"}, + want: "198.51.100.17", + }, + { + name: "forwarded header quoted", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{"Forwarded": `for="198.51.100.17"`}, + want: "198.51.100.17", + }, + { + name: "forwarded header ipv6 bracketed", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{"Forwarded": `for="[2001:db8::1]"`}, + want: "2001:db8::1", + }, + { + name: "forwarded header with port", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{"Forwarded": `for="198.51.100.17:4711"`}, + want: "198.51.100.17", + }, + { + name: "forwarded header chain", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{"Forwarded": "for=198.51.100.17, for=70.41.3.18"}, + want: "198.51.100.17", + }, + { + name: "forwarded takes precedence over x-forwarded-for", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{ + "Forwarded": "for=198.51.100.17", + "X-Forwarded-For": "203.0.113.50", + }, + want: "198.51.100.17", + }, + { + name: "forwarded with extra directives", + remoteAddr: "10.0.0.1:1234", + headers: map[string]string{"Forwarded": "for=198.51.100.17;proto=https;by=203.0.113.60"}, + want: "198.51.100.17", + }, + } + + for _, tt := range tests { + t.Run( + tt.name, + func(t *testing.T) { + t.Parallel() + + r := httptest.NewRequest("GET", "/", nil) + r.RemoteAddr = tt.remoteAddr + for k, v := range tt.headers { + r.Header.Set(k, v) + } + + assert.Equal(t, tt.want, clientip.Extract(r)) + }, + ) + } +} diff --git a/pkg/server/api/cookiebanner/v1/handler.go b/pkg/server/api/cookiebanner/v1/handler.go index 7a6593b54..f3262f4c1 100644 --- a/pkg/server/api/cookiebanner/v1/handler.go +++ b/pkg/server/api/cookiebanner/v1/handler.go @@ -47,7 +47,6 @@ func NewMux( r := chi.NewMux() r.Use(newCORSMiddleware(logger, cookieBannerSvc)) - r.Use(clientip.NewMiddleware()) r.Get("/{bannerID}/config", h.handleGetConfig) r.Get("/{bannerID}/consents/{visitorID}", h.handleGetConsent) r.Post("/{bannerID}/consents", h.handlePostConsent) @@ -140,7 +139,7 @@ func (h *Handler) handlePostConsent(w http.ResponseWriter, r *http.Request) { return } - ip := clientip.FromContext(r.Context()) + ip := clientip.Extract(r) ua := r.UserAgent() req := cookiebanner.RecordConsentRequest{ diff --git a/pkg/server/trustedproxy/trustedproxy.go b/pkg/server/trustedproxy/trustedproxy.go new file mode 100644 index 000000000..1d3cf4bc8 --- /dev/null +++ b/pkg/server/trustedproxy/trustedproxy.go @@ -0,0 +1,62 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package trustedproxy + +import ( + "net" + "net/http" +) + +var forwardedHeaders = []string{ + "Forwarded", + "X-Forwarded-For", +} + +// NewMiddleware returns an HTTP middleware that strips forwarded +// headers from requests that did not originate from one of the given +// trusted proxy IPs. When the list is empty every request is treated +// as untrusted and the headers are always removed. +func NewMiddleware(trusted []net.IP) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if !isTrusted(r.RemoteAddr, trusted) { + for _, h := range forwardedHeaders { + r.Header.Del(h) + } + } + next.ServeHTTP(w, r) + }) + } +} + +func isTrusted(remoteAddr string, trusted []net.IP) bool { + host, _, err := net.SplitHostPort(remoteAddr) + if err != nil { + host = remoteAddr + } + + ip := net.ParseIP(host) + if ip == nil { + return false + } + + for _, t := range trusted { + if t.Equal(ip) { + return true + } + } + + return false +} diff --git a/pkg/server/trustedproxy/trustedproxy_test.go b/pkg/server/trustedproxy/trustedproxy_test.go new file mode 100644 index 000000000..9b831d3bf --- /dev/null +++ b/pkg/server/trustedproxy/trustedproxy_test.go @@ -0,0 +1,137 @@ +// Copyright (c) 2026 Probo Inc . +// +// Permission to use, copy, modify, and/or distribute this software for any +// purpose with or without fee is hereby granted, provided that the above +// copyright notice and this permission notice appear in all copies. +// +// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH +// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY +// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT, +// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM +// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR +// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR +// PERFORMANCE OF THIS SOFTWARE. + +package trustedproxy_test + +import ( + "net" + "net/http" + "net/http/httptest" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "go.probo.inc/probo/pkg/server/trustedproxy" +) + +func newRequest(remoteAddr string, headers map[string]string) *http.Request { + r := httptest.NewRequest(http.MethodGet, "/", nil) + r.RemoteAddr = remoteAddr + for k, v := range headers { + r.Header.Set(k, v) + } + return r +} + +func TestNewMiddleware(t *testing.T) { + t.Parallel() + + t.Run( + "strips forwarded headers from untrusted proxy", + func(t *testing.T) { + t.Parallel() + + trusted := []net.IP{net.ParseIP("10.0.0.1")} + middleware := trustedproxy.NewMiddleware(trusted) + + var captured *http.Request + handler := middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + captured = r + })) + + r := newRequest("192.168.1.1:1234", map[string]string{ + "X-Forwarded-For": "203.0.113.50", + "Forwarded": "for=198.51.100.17", + }) + handler.ServeHTTP(httptest.NewRecorder(), r) + + require.NotNil(t, captured) + assert.Empty(t, captured.Header.Get("X-Forwarded-For")) + assert.Empty(t, captured.Header.Get("Forwarded")) + }, + ) + + t.Run( + "preserves forwarded headers from trusted proxy", + func(t *testing.T) { + t.Parallel() + + trusted := []net.IP{net.ParseIP("10.0.0.1")} + middleware := trustedproxy.NewMiddleware(trusted) + + var captured *http.Request + handler := middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + captured = r + })) + + r := newRequest("10.0.0.1:1234", map[string]string{ + "X-Forwarded-For": "203.0.113.50", + "Forwarded": "for=198.51.100.17", + }) + handler.ServeHTTP(httptest.NewRecorder(), r) + + require.NotNil(t, captured) + assert.Equal(t, "203.0.113.50", captured.Header.Get("X-Forwarded-For")) + assert.Equal(t, "for=198.51.100.17", captured.Header.Get("Forwarded")) + }, + ) + + t.Run( + "empty trusted list strips all forwarded headers", + func(t *testing.T) { + t.Parallel() + + middleware := trustedproxy.NewMiddleware(nil) + + var captured *http.Request + handler := middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + captured = r + })) + + r := newRequest("10.0.0.1:1234", map[string]string{ + "X-Forwarded-For": "203.0.113.50", + }) + handler.ServeHTTP(httptest.NewRecorder(), r) + + require.NotNil(t, captured) + assert.Empty(t, captured.Header.Get("X-Forwarded-For")) + }, + ) + + t.Run( + "multiple trusted proxies", + func(t *testing.T) { + t.Parallel() + + trusted := []net.IP{ + net.ParseIP("10.0.0.1"), + net.ParseIP("10.0.0.2"), + } + middleware := trustedproxy.NewMiddleware(trusted) + + var captured *http.Request + handler := middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + captured = r + })) + + r := newRequest("10.0.0.2:5678", map[string]string{ + "X-Forwarded-For": "203.0.113.50", + }) + handler.ServeHTTP(httptest.NewRecorder(), r) + + require.NotNil(t, captured) + assert.Equal(t, "203.0.113.50", captured.Header.Get("X-Forwarded-For")) + }, + ) +}