mirror of
https://github.com/oauth2-proxy/oauth2-proxy.git
synced 2026-10-06 14:41:22 +02:00
@@ -2,9 +2,12 @@ package middleware
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/apis/sessions"
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/ip"
|
||||
)
|
||||
|
||||
type scopeKey string
|
||||
@@ -18,9 +21,13 @@ const RequestScopeKey scopeKey = "request-scope"
|
||||
// within the chain.
|
||||
type RequestScope struct {
|
||||
// ReverseProxy tracks whether OAuth2-Proxy is operating in reverse proxy
|
||||
// mode and if request `X-Forwarded-*` headers should be trusted
|
||||
// mode and if request `X-Forwarded-*` headers may be trusted
|
||||
ReverseProxy bool
|
||||
|
||||
// TrustedProxies tracks which direct callers are allowed to supply
|
||||
// forwarded headers when ReverseProxy mode is enabled.
|
||||
TrustedProxies *ip.NetSet
|
||||
|
||||
// RequestID is set to the request's `X-Request-Id` header if set.
|
||||
// Otherwise a random UUID is set.
|
||||
RequestID string
|
||||
@@ -58,3 +65,43 @@ func AddRequestScope(req *http.Request, scope *RequestScope) *http.Request {
|
||||
ctx := context.WithValue(req.Context(), RequestScopeKey, scope)
|
||||
return req.WithContext(ctx)
|
||||
}
|
||||
|
||||
// CanTrustForwardedHeaders returns whether forwarded headers should be
|
||||
// processed for this request.
|
||||
func (s *RequestScope) CanTrustForwardedHeaders(req *http.Request) bool {
|
||||
if s == nil || req == nil || !s.ReverseProxy || s.TrustedProxies == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
if isUnixSocketRemoteAddr(req.RemoteAddr) {
|
||||
return true
|
||||
}
|
||||
|
||||
remoteIP := parseRemoteAddrIP(req.RemoteAddr)
|
||||
if remoteIP == nil {
|
||||
return false
|
||||
}
|
||||
|
||||
return s.TrustedProxies.Has(remoteIP)
|
||||
}
|
||||
|
||||
func parseRemoteAddrIP(remoteAddr string) net.IP {
|
||||
if remoteAddr == "" {
|
||||
return nil
|
||||
}
|
||||
|
||||
if ip := net.ParseIP(remoteAddr); ip != nil {
|
||||
return ip
|
||||
}
|
||||
|
||||
host, _, err := net.SplitHostPort(remoteAddr)
|
||||
if err != nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
return net.ParseIP(host)
|
||||
}
|
||||
|
||||
func isUnixSocketRemoteAddr(remoteAddr string) bool {
|
||||
return remoteAddr == "@" || strings.HasPrefix(remoteAddr, "/")
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"net/http"
|
||||
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/apis/middleware"
|
||||
"github.com/oauth2-proxy/oauth2-proxy/v7/pkg/ip"
|
||||
. "github.com/onsi/ginkgo/v2"
|
||||
. "github.com/onsi/gomega"
|
||||
)
|
||||
@@ -53,4 +54,37 @@ var _ = Describe("Scope Suite", func() {
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
Context("CanTrustForwardedHeaders", func() {
|
||||
var request *http.Request
|
||||
var scope *middleware.RequestScope
|
||||
|
||||
BeforeEach(func() {
|
||||
var err error
|
||||
request, err = http.NewRequest("", "http://127.0.0.1/", nil)
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
|
||||
trustedProxies, err := ip.ParseNetSet([]string{"127.0.0.1"})
|
||||
Expect(err).ToNot(HaveOccurred())
|
||||
scope = &middleware.RequestScope{
|
||||
ReverseProxy: true,
|
||||
TrustedProxies: trustedProxies,
|
||||
}
|
||||
})
|
||||
|
||||
It("returns true for a trusted remote address", func() {
|
||||
request.RemoteAddr = "127.0.0.1:4180"
|
||||
Expect(scope.CanTrustForwardedHeaders(request)).To(BeTrue())
|
||||
})
|
||||
|
||||
It("returns false for an untrusted remote address", func() {
|
||||
request.RemoteAddr = "192.0.2.10:4180"
|
||||
Expect(scope.CanTrustForwardedHeaders(request)).To(BeFalse())
|
||||
})
|
||||
|
||||
It("returns true for unix socket callers", func() {
|
||||
request.RemoteAddr = "@"
|
||||
Expect(scope.CanTrustForwardedHeaders(request)).To(BeTrue())
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user