Merge branch 'reverse-proxy-allow-match-or' into reverse-proxy-crowdsec-appsec

# Conflicts:
#	proxy/internal/auth/middleware.go
#	proxy/internal/auth/middleware_test.go
#	proxy/internal/auth/tunnel_lookup_test.go
#	proxy/management_integration_test.go
#	proxy/server.go
#	shared/management/proto/proxy_service.pb.go
This commit is contained in:
Viktor Liu
2026-09-23 07:53:31 +02:00
968 changed files with 72582 additions and 16724 deletions
+26
View File
@@ -0,0 +1,26 @@
# PIN and password authentication limits
PIN and password credentials are accepted only in a POST form body. Query-string
credentials and credentials on other HTTP methods are ignored.
The proxy permits a burst of five credential checks per account and service,
then replenishes one check every six seconds (ten per minute). PIN and password
checks share the same budget. Five failed checks from one client IP in a
rolling five-minute window block that source for fifteen minutes. In-flight checks
reserve failure slots; blocked requests do not extend the cooldown. Successful
authentication clears that source's failure history. Infrastructure failures
consume the service budget without counting as incorrect credentials.
Throttled requests return HTTP 429 with a `Retry-After` delay in seconds. The
login page displays that delay. Existing authenticated sessions and other
authentication methods do not consume these credential budgets.
The client IP comes from the existing trusted-proxy resolution. Deployments
behind a load balancer must configure trusted proxies correctly; otherwise
visitors share the load balancer's source budget. Visitors behind the same NAT
also share a source budget for a service.
State is held in memory per proxy process and resets on restart. Multiple
replicas have independent budgets. State is bounded to 16,384 source entries and
4,096 service entries; when capacity is exhausted, new checks are denied until
idle entries expire. Active blocks are never evicted to admit a new source.
+100
View File
@@ -0,0 +1,100 @@
package auth
import (
"errors"
"math"
"net/http"
"strconv"
"time"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/proxy/internal/proxy"
)
var errCredentialClientIP = errors.New("invalid client address")
type credentialLimitError struct {
retryAfter time.Duration
}
func (e *credentialLimitError) Error() string {
return "too many authentication attempts"
}
func credentialFormValue(r *http.Request, field string) string {
if r.Method != http.MethodPost {
return ""
}
return r.PostFormValue(field)
}
func (mw *Middleware) authenticateScheme(r *http.Request, config DomainConfig, scheme Scheme) (string, string, error) {
method := scheme.Type()
if (method != auth.MethodPIN && method != auth.MethodPassword) || !wasCredentialSubmitted(r, method) {
return scheme.Authenticate(r)
}
ip := mw.resolveClientIP(r).Unmap()
if !ip.IsValid() {
return "", "", errCredentialClientIP
}
source, retry := mw.credentials.begin(credentialSourceKey{
service: credentialServiceKey{accountID: config.AccountID, serviceID: config.ServiceID},
ip: ip,
})
if retry > 0 {
return "", "", &credentialLimitError{retryAfter: retry}
}
token, prompt, err := scheme.Authenticate(r)
outcome := credentialUnavailable
if err == nil {
outcome = credentialRejected
if token != "" {
outcome = credentialAccepted
}
}
mw.credentials.finish(source, outcome)
return token, prompt, err
}
func credentialRetryAfter(err error) time.Duration {
var limitErr *credentialLimitError
if errors.As(err, &limitErr) {
return limitErr.retryAfter
}
s := status.Convert(err)
if s.Code() != codes.ResourceExhausted {
return 0
}
for _, detail := range s.Details() {
if info, ok := detail.(*errdetails.RetryInfo); ok && info.RetryDelay != nil && info.RetryDelay.CheckValid() == nil {
if delay := info.RetryDelay.AsDuration(); delay > 0 {
return delay
}
}
}
return credentialCheckInterval
}
func (mw *Middleware) writeAuthenticationError(w http.ResponseWriter, r *http.Request, method auth.Method, err error) {
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
cd.SetOrigin(proxy.OriginAuth)
cd.SetAuthMethod(method.String())
}
if retry := credentialRetryAfter(err); retry > 0 {
// RFC 6585 section 4 forbids caching 429 responses.
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Retry-After", strconv.FormatInt(int64(math.Ceil(retry.Seconds())), 10))
http.Error(w, "too many authentication attempts; try again later", http.StatusTooManyRequests)
return
}
if errors.Is(err, errCredentialClientIP) {
http.Error(w, "invalid client address", http.StatusBadRequest)
return
}
mw.logger.WithField("scheme", method.String()).Warnf("authentication infrastructure error: %v", err)
http.Error(w, "authentication service unavailable", http.StatusBadGateway)
}
+169
View File
@@ -0,0 +1,169 @@
package auth
import (
"net/netip"
"sync"
"time"
"golang.org/x/time/rate"
"github.com/netbirdio/netbird/proxy/internal/types"
)
const (
credentialFailureLimit = 5
credentialFailureWindow = 5 * time.Minute
credentialBlockDuration = 15 * time.Minute
credentialCheckInterval = 6 * time.Second
credentialCheckBurst = 5
credentialMaxSources = 16384
credentialMaxServices = 4096
credentialCleanupInterval = time.Minute
)
type credentialServiceKey struct {
accountID types.AccountID
serviceID types.ServiceID
}
type credentialSourceKey struct {
service credentialServiceKey
ip netip.Addr
}
type credentialSource struct {
failures []time.Time
pending int
expiresAt time.Time
blockedUntil time.Time
}
type credentialService struct {
limiter *rate.Limiter
lastUsed time.Time
}
type credentialOutcome string
const (
credentialUnavailable credentialOutcome = "unavailable"
credentialRejected credentialOutcome = "rejected"
credentialAccepted credentialOutcome = "accepted"
)
// State is local to this proxy process. Active blocks are never evicted to
// make room for a new source; exhausting capacity denies new checks.
type credentialLimiter struct {
mu sync.Mutex
now func() time.Time
sources map[credentialSourceKey]*credentialSource
services map[credentialServiceKey]*credentialService
nextCleanup time.Time
}
func newCredentialLimiter() *credentialLimiter {
return &credentialLimiter{
now: time.Now,
sources: make(map[credentialSourceKey]*credentialSource),
services: make(map[credentialServiceKey]*credentialService),
}
}
func (l *credentialLimiter) begin(key credentialSourceKey) (*credentialSource, time.Duration) {
l.mu.Lock()
defer l.mu.Unlock()
now := l.now()
l.cleanup(now)
source := l.sources[key]
if source != nil {
if now.Before(source.blockedUntil) {
return nil, source.blockedUntil.Sub(now)
}
if source.pending == 0 && !now.Before(source.expiresAt) {
*source = credentialSource{}
}
source.expireFailures(now)
// Reserve the failure budget before verification so concurrent guesses
// cannot all pass a check against the same completed failure count.
if len(source.failures)+source.pending >= credentialFailureLimit {
return nil, time.Second
}
} else if len(l.sources) >= credentialMaxSources {
return nil, credentialCleanupInterval
}
if retry := l.allowService(key.service, now); retry > 0 {
return nil, retry
}
if source == nil {
source = &credentialSource{}
l.sources[key] = source
}
if source.expiresAt.IsZero() {
source.expiresAt = now.Add(credentialFailureWindow)
}
source.pending++
return source, 0
}
func (l *credentialLimiter) allowService(key credentialServiceKey, now time.Time) time.Duration {
service := l.services[key]
if service == nil {
if len(l.services) >= credentialMaxServices {
return credentialCleanupInterval
}
service = &credentialService{limiter: rate.NewLimiter(rate.Every(credentialCheckInterval), credentialCheckBurst)}
l.services[key] = service
}
service.lastUsed = now
if service.limiter.AllowN(now, 1) {
return 0
}
return max(time.Nanosecond, time.Duration((1-service.limiter.TokensAt(now))*float64(credentialCheckInterval)))
}
func (l *credentialLimiter) finish(source *credentialSource, outcome credentialOutcome) {
l.mu.Lock()
defer l.mu.Unlock()
source.pending--
now := l.now()
source.expireFailures(now)
switch outcome {
case credentialRejected:
source.failures = append(source.failures, now)
source.expiresAt = now.Add(credentialFailureWindow)
if len(source.failures) >= credentialFailureLimit && source.blockedUntil.IsZero() {
source.blockedUntil = now.Add(credentialBlockDuration)
source.expiresAt = source.blockedUntil
}
case credentialAccepted:
if !now.Before(source.blockedUntil) {
source.failures = nil
source.expiresAt = now.Add(credentialFailureWindow)
}
case credentialUnavailable:
// Transport failures consume the service budget, but are not bad guesses.
}
}
func (s *credentialSource) expireFailures(now time.Time) {
for len(s.failures) > 0 && !now.Before(s.failures[0].Add(credentialFailureWindow)) {
s.failures = s.failures[1:]
}
}
func (l *credentialLimiter) cleanup(now time.Time) {
if now.Before(l.nextCleanup) {
return
}
l.nextCleanup = now.Add(credentialCleanupInterval)
for key, source := range l.sources {
if source.pending == 0 && !now.Before(source.expiresAt) {
delete(l.sources, key)
}
}
for key, service := range l.services {
if now.Sub(service.lastUsed) >= credentialBlockDuration {
delete(l.services, key)
}
}
}
@@ -0,0 +1,190 @@
package auth
import (
"net/netip"
"strconv"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/proxy/internal/types"
)
func TestCredentialLimiterCooldown(t *testing.T) {
l := newCredentialLimiter()
now := time.Now()
l.now = func() time.Time { return now }
key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")}
for range credentialFailureLimit {
attempt, retry := l.begin(key)
require.Zero(t, retry, "initial guesses must reach verification")
l.finish(attempt, credentialRejected)
}
_, retry := l.begin(key)
assert.Equal(t, credentialBlockDuration, retry, "five failures must start a fifteen-minute block")
now = now.Add(credentialBlockDuration - time.Second)
_, retry = l.begin(key)
assert.Equal(t, time.Second, retry, "blocked requests must not extend the deadline")
now = now.Add(time.Second)
attempt, retry := l.begin(key)
require.Zero(t, retry, "the source must recover when its block expires")
l.finish(attempt, credentialAccepted)
}
func TestCredentialLimiterFailureWindowAndSuccess(t *testing.T) {
for _, outcome := range []credentialOutcome{credentialAccepted, credentialUnavailable} {
t.Run(map[credentialOutcome]string{credentialAccepted: "success", credentialUnavailable: "infrastructure error"}[outcome], func(t *testing.T) {
l := newCredentialLimiter()
now := time.Now()
l.now = func() time.Time { return now }
key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")}
for range 4 {
attempt, retry := l.begin(key)
require.Zero(t, retry, "four failures must fit the budget")
l.finish(attempt, credentialRejected)
}
attempt, retry := l.begin(key)
require.Zero(t, retry, "fifth check must be allowed")
l.finish(attempt, outcome)
now = now.Add(credentialCheckInterval)
attempt, retry = l.begin(key)
require.Zero(t, retry, "success or infrastructure error must not start a block")
l.finish(attempt, credentialRejected)
now = now.Add(credentialCheckInterval)
attempt, retry = l.begin(key)
if outcome == credentialUnavailable {
assert.Greater(t, retry, time.Duration(0), "infrastructure errors must preserve earlier failures")
return
}
require.Zero(t, retry, "success must clear earlier failures")
l.finish(attempt, credentialRejected)
now = now.Add(credentialFailureWindow)
for range credentialFailureLimit {
attempt, retry = l.begin(key)
require.Zero(t, retry, "old failures must expire")
l.finish(attempt, credentialRejected)
}
})
}
}
func TestCredentialLimiterRollingWindow(t *testing.T) {
l := newCredentialLimiter()
now := time.Now()
l.now = func() time.Time { return now }
key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")}
attempt, retry := l.begin(key)
require.Zero(t, retry, "the first failure starts the history")
l.finish(attempt, credentialRejected)
now = now.Add(4 * time.Minute)
for range 3 {
attempt, retry = l.begin(key)
require.Zero(t, retry, "three more failures must fit the budget")
l.finish(attempt, credentialRejected)
}
now = now.Add(time.Minute + time.Second)
for range 2 {
attempt, retry = l.begin(key)
require.Zero(t, retry, "only the oldest failure must have expired")
l.finish(attempt, credentialRejected)
}
_, retry = l.begin(key)
assert.Equal(t, credentialBlockDuration, retry, "five recent failures must block even across the first window boundary")
}
func TestCredentialLimiterServiceBudget(t *testing.T) {
l := newCredentialLimiter()
now := time.Now()
l.now = func() time.Time { return now }
key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")}
for range credentialCheckBurst {
attempt, retry := l.begin(key)
require.Zero(t, retry, "initial checks must fit the service burst")
l.finish(attempt, credentialAccepted)
key.ip = key.ip.Next()
}
_, retry := l.begin(key)
assert.Equal(t, credentialCheckInterval, retry, "changing IP must not bypass the service budget")
other := key
other.service.accountID = "another-account"
attempt, retry := l.begin(other)
require.Zero(t, retry, "accounts must have separate budgets")
l.finish(attempt, credentialAccepted)
other = key
other.service.serviceID = "another-service"
attempt, retry = l.begin(other)
require.Zero(t, retry, "services must have separate budgets")
l.finish(attempt, credentialAccepted)
now = now.Add(credentialCheckInterval)
attempt, retry = l.begin(key)
require.Zero(t, retry, "one check must refill every six seconds")
l.finish(attempt, credentialAccepted)
_, retry = l.begin(key)
assert.Equal(t, credentialCheckInterval, retry, "refill must only grant one new check")
}
func TestCredentialLimiterConcurrentReservations(t *testing.T) {
l := newCredentialLimiter()
now := time.Now()
l.now = func() time.Time { return now }
key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")}
var attempts []*credentialSource
for range credentialFailureLimit {
attempt, retry := l.begin(key)
require.Zero(t, retry, "initial requests must reserve the failure budget")
attempts = append(attempts, attempt)
}
// Refill the service budget while earlier verification calls are still running.
now = now.Add(time.Minute)
var admitted atomic.Int32
var wg sync.WaitGroup
for range 100 {
wg.Go(func() {
attempt, retry := l.begin(key)
if retry == 0 {
admitted.Add(1)
l.finish(attempt, credentialRejected)
}
})
}
wg.Wait()
assert.Zero(t, admitted.Load(), "in-flight guesses must reserve the failure budget despite a refilled service budget")
for _, attempt := range attempts {
wg.Go(func() { l.finish(attempt, credentialRejected) })
}
wg.Wait()
_, retry := l.begin(key)
assert.Equal(t, credentialBlockDuration, retry, "concurrent failures must activate the block")
}
func TestCredentialLimiterCapacityAndCleanup(t *testing.T) {
for _, fullSources := range []bool{true, false} {
t.Run(map[bool]string{true: "sources", false: "services"}[fullSources], func(t *testing.T) {
l := newCredentialLimiter()
now := time.Now()
l.now = func() time.Time { return now }
key := credentialSourceKey{service: credentialServiceKey{"account", "service"}, ip: netip.MustParseAddr("192.0.2.1")}
if fullSources {
ip := netip.MustParseAddr("198.18.0.1")
for range credentialMaxSources {
l.sources[credentialSourceKey{service: key.service, ip: ip}] = &credentialSource{expiresAt: now.Add(credentialBlockDuration), blockedUntil: now.Add(credentialBlockDuration)}
ip = ip.Next()
}
} else {
for i := range credentialMaxServices {
l.services[credentialServiceKey{serviceID: key.service.serviceID, accountID: types.AccountID(strconv.Itoa(i))}] = &credentialService{lastUsed: now}
}
}
_, retry := l.begin(key)
assert.Positive(t, retry, "full state must deny new checks without evicting active entries")
now = now.Add(credentialBlockDuration)
attempt, retry := l.begin(key)
require.Zero(t, retry, "expired state must release capacity")
l.finish(attempt, credentialAccepted)
})
}
}
+196
View File
@@ -0,0 +1,196 @@
package auth
import (
"context"
"fmt"
"net/http"
"net/http/httptest"
"net/netip"
"net/url"
"strings"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/types/known/durationpb"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
servicemanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service/manager"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/store"
mgmttypes "github.com/netbirdio/netbird/management/server/types"
proxyauth "github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/proxy/internal/proxy"
"github.com/netbirdio/netbird/shared/management/proto"
)
// localCredentialClient replaces the transport while keeping the real service
// store, credential verification, and session signing.
type localCredentialClient struct {
server *nbgrpc.ProxyServiceServer
}
func (c localCredentialClient) Authenticate(ctx context.Context, req *proto.AuthenticateRequest, _ ...grpc.CallOption) (*proto.AuthenticateResponse, error) {
return c.server.Authenticate(ctx, req)
}
func credentialHandler(t *testing.T, field string) (*Middleware, http.Handler) {
t.Helper()
ctx := context.Background()
s, err := store.NewStore(ctx, mgmttypes.SqliteStoreEngine, t.TempDir(), nil, false)
require.NoError(t, err)
t.Cleanup(func() { assert.NoError(t, s.Close(ctx)) })
require.NoError(t, s.SaveAccount(ctx, &mgmttypes.Account{Id: "account"}))
keys := generateTestKeyPair(t)
svc := &service.Service{
ID: "service", AccountID: "account", Name: "test", Domain: "example.com",
Enabled: true, SessionPrivateKey: keys.PrivateKey, SessionPublicKey: keys.PublicKey,
Auth: service.AuthConfig{
PinAuth: &service.PINAuthConfig{Enabled: true, Pin: "842716"},
PasswordAuth: &service.PasswordAuthConfig{Enabled: true, Password: "842716"},
},
}
require.NoError(t, svc.Auth.HashSecrets())
require.NoError(t, s.CreateService(ctx, svc))
server := nbgrpc.NewProxyServiceServer(nil, nil, nil, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
t.Cleanup(server.Close)
server.SetServiceManager(servicemanager.NewManager(s, nil, nil, nil, nil, nil))
client := localCredentialClient{server: server}
var scheme Scheme = NewPin(client, "service", "account")
if field == "password" {
scheme = NewPassword(client, "service", "account")
}
mw := NewMiddleware(nil, nil, nil)
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: keys.PublicKey, SessionExpiration: time.Hour, AccountID: "account", ServiceID: "service"}))
return mw, mw.Protect(newPassthroughHandler())
}
func credentialRequest(method, field, value string) *http.Request {
r := httptest.NewRequest(method, "https://example.com/", strings.NewReader(url.Values{field: {value}}.Encode()))
r.Header.Set("Content-Type", "application/x-www-form-urlencoded")
r.RemoteAddr = "198.51.100.25:12345"
return r
}
func TestCredentialAuthPOSTOnly(t *testing.T) {
for _, field := range []string{"pin", "password"} {
t.Run(field, func(t *testing.T) {
_, handler := credentialHandler(t, field)
for _, method := range []string{http.MethodGet, http.MethodPut, http.MethodPatch, http.MethodDelete, http.MethodPost} {
r := credentialRequest(method, field, "")
r.URL.RawQuery = url.Values{field: {"842716"}}.Encode()
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, r)
assert.Equal(t, http.StatusUnauthorized, resp.Code, "%s query credentials must not authenticate", method)
assert.Empty(t, resp.Result().Cookies(), "query credentials must not issue a session")
}
for _, method := range []string{http.MethodGet, http.MethodPut, http.MethodPatch, http.MethodDelete} {
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, credentialRequest(method, field, "842716"))
assert.Equal(t, http.StatusUnauthorized, resp.Code, "%s body credentials must not authenticate", method)
}
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, credentialRequest(http.MethodPost, field, "842716"))
assert.Equal(t, http.StatusSeeOther, resp.Code, "POST body credentials must authenticate")
})
}
}
func TestCredentialAuthThrottling(t *testing.T) {
for _, field := range []string{"pin", "password"} {
t.Run(field, func(t *testing.T) {
_, handler := credentialHandler(t, field)
for range 5 {
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, credentialRequest(http.MethodPost, field, "000000"))
require.Equal(t, http.StatusUnauthorized, resp.Code, "initial wrong credentials must be rejected")
}
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, credentialRequest(http.MethodPost, field, "842716"))
assert.Equal(t, http.StatusTooManyRequests, resp.Code, "even correct credentials must wait for the block to expire")
assert.Equal(t, "900", resp.Header().Get("Retry-After"), "five failures must block the source for fifteen minutes")
assert.Empty(t, resp.Result().Cookies(), "blocked credentials must not issue a session")
})
}
}
func TestCredentialAuthSessionAndClientIP(t *testing.T) {
keys := generateTestKeyPair(t)
token, err := sessionkey.SignToken(keys.PrivateKey, "pin-user", "", "example.com", proxyauth.MethodPIN, nil, nil, time.Hour)
require.NoError(t, err)
mw := NewMiddleware(nil, nil, nil)
now := time.Now()
mw.credentials.now = func() time.Time { return now }
scheme := &stubScheme{method: proxyauth.MethodPIN, promptID: "pin"}
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: keys.PublicKey, AccountID: "account", ServiceID: "service"}))
handler := mw.Protect(newPassthroughHandler())
for range credentialFailureLimit {
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, credentialRequest(http.MethodPost, "pin", "000000"))
require.Equal(t, http.StatusUnauthorized, resp.Code, "bad PIN must consume the failure budget")
}
now = now.Add(credentialCheckInterval)
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: keys.PublicKey, AccountID: "account", ServiceID: "service"}))
r := credentialRequest(http.MethodPost, "pin", "000000")
r.RemoteAddr = "[::ffff:198.51.100.25]:45678"
r.Header.Set("X-Forwarded-For", "192.0.2.5")
r.Header.Set("X-Real-IP", "192.0.2.6")
resp := httptest.NewRecorder()
handler.ServeHTTP(resp, r)
assert.Equal(t, http.StatusTooManyRequests, resp.Code, "mapped addresses and untrusted forwarding headers must not bypass the source block")
assert.Equal(t, "no-store", resp.Header().Get("Cache-Control"), "rate limits must not be cached")
r.AddCookie(&http.Cookie{Name: proxyauth.SessionCookieName, Value: token})
resp = httptest.NewRecorder()
handler.ServeHTTP(resp, r)
assert.Equal(t, http.StatusOK, resp.Code, "an existing session must pass even with credentials in the request")
assert.Equal(t, "backend", resp.Body.String(), "the authenticated request must reach the application")
r = credentialRequest(http.MethodPost, "pin", "000000")
cd := proxy.NewCapturedData("test")
cd.SetClientIP(netip.MustParseAddr("192.0.2.9"))
r = r.WithContext(proxy.WithCapturedData(r.Context(), cd))
resp = httptest.NewRecorder()
handler.ServeHTTP(resp, r)
assert.Equal(t, http.StatusUnauthorized, resp.Code, "a client resolved by the trusted-proxy middleware must get its own source budget")
r = credentialRequest(http.MethodPost, "pin", "000000")
r.RemoteAddr = "invalid"
resp = httptest.NewRecorder()
handler.ServeHTTP(resp, r)
assert.Equal(t, http.StatusBadRequest, resp.Code, "an unresolvable client address must fail closed")
now = now.Add(credentialBlockDuration)
scheme.token = token
resp = httptest.NewRecorder()
handler.ServeHTTP(resp, credentialRequest(http.MethodPost, "pin", "842716"))
assert.Equal(t, http.StatusSeeOther, resp.Code, "credentials must work again after cooldown")
}
func TestCredentialAuthManagementThrottling(t *testing.T) {
s, err := status.New(codes.ResourceExhausted, "rate limited").WithDetails(&errdetails.RetryInfo{RetryDelay: durationpb.New(2500 * time.Millisecond)})
require.NoError(t, err)
for _, tc := range []struct {
name string
err error
code int
retry string
}{
{"retry info", fmt.Errorf("authenticate PIN: %w", s.Err()), http.StatusTooManyRequests, "3"},
{"missing retry info", status.Error(codes.ResourceExhausted, "rate limited"), http.StatusTooManyRequests, "6"},
{"unavailable", status.Error(codes.Unavailable, "unavailable"), http.StatusBadGateway, ""},
} {
t.Run(tc.name, func(t *testing.T) {
keys := generateTestKeyPair(t)
mw := NewMiddleware(nil, nil, nil)
scheme := &stubScheme{method: proxyauth.MethodPIN, authFn: func(*http.Request) (string, string, error) { return "", "", tc.err }}
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: keys.PublicKey, AccountID: "account", ServiceID: "service"}))
resp := httptest.NewRecorder()
mw.Protect(newPassthroughHandler()).ServeHTTP(resp, credentialRequest(http.MethodPost, "pin", "000000"))
assert.Equal(t, tc.code, resp.Code, "management errors must keep their HTTP meaning")
assert.Equal(t, tc.retry, resp.Header().Get("Retry-After"), "retry hints must round up to whole seconds")
})
}
}
+94 -12
View File
@@ -73,6 +73,11 @@ type DomainConfig struct {
redactBodyFields []string
redactHeaders []string
redactQueryParams []string
// AllowedGroups holds the group ids that may reach the service through an
// OIDC identity. When non-empty, a session cookie is honoured only if its
// groups claim intersects this set. Empty means group membership does not
// restrict access.
AllowedGroups map[string]struct{}
}
type validationResult struct {
@@ -96,6 +101,7 @@ type Middleware struct {
sessionValidator SessionValidator
geo restrict.GeoResolver
tunnelCache *tunnelValidationCache
credentials *credentialLimiter
// appsec is the shared CrowdSec AppSec client, nil when the proxy has no
// AppSec endpoint configured. Set once during startup, before serving.
appsec *appsec.Client
@@ -113,6 +119,7 @@ func NewMiddleware(logger *log.Logger, sessionValidator SessionValidator, geo re
sessionValidator: sessionValidator,
geo: geo,
tunnelCache: newTunnelValidationCache(),
credentials: newCredentialLimiter(),
}
}
@@ -159,7 +166,7 @@ func (mw *Middleware) Protect(next http.Handler) http.Handler {
if mw.forwardWithTunnelPeer(w, r, host, config, next) {
return
}
http.Error(w, "Forbidden", http.StatusForbidden)
denyPrivate(w)
return
}
@@ -254,7 +261,7 @@ func (mw *Middleware) checkIPRestrictions(w http.ResponseWriter, r *http.Request
clientIP := mw.resolveClientIP(r)
if !clientIP.IsValid() {
mw.logger.Debugf("IP restriction: cannot resolve client address for %q, denying", r.RemoteAddr)
http.Error(w, "Forbidden", http.StatusForbidden)
denyForbidden(w, config)
return false
}
@@ -289,7 +296,7 @@ func (mw *Middleware) checkIPRestrictions(w http.ResponseWriter, r *http.Request
reason := verdict.String()
mw.blockIPRestriction(r, reason)
http.Error(w, "Forbidden", http.StatusForbidden)
denyForbidden(w, config)
return false
}
@@ -333,7 +340,7 @@ func (mw *Middleware) checkAppSec(w http.ResponseWriter, r *http.Request, config
mw.markDenied(r, verdict.String())
mw.logger.Debugf("AppSec: %s for %s %s", verdict, r.Host, r.RemoteAddr)
http.Error(w, "Forbidden", http.StatusForbidden)
denyForbidden(w, config)
return false, release
}
@@ -421,6 +428,26 @@ func credentialHeaders(schemes []Scheme) []string {
return names
}
// denyForbidden writes a 403, dropping the client connection when the
// domain is private so a later retry cannot reuse it.
func denyForbidden(w http.ResponseWriter, config DomainConfig) {
if config.Private {
denyPrivate(w)
return
}
http.Error(w, "Forbidden", http.StatusForbidden)
}
// denyPrivate writes a 403 and closes the connection, so a client refused
// before joining the overlay cannot keep retrying on the same warm socket.
// Go's HTTP/2 server turns the exact lowercase "close" token into a GOAWAY.
func denyPrivate(w http.ResponseWriter) {
h := w.Header()
h.Set("Connection", "close")
h.Set("Cache-Control", "no-store")
http.Error(w, "Forbidden", http.StatusForbidden)
}
// resolveClientIP extracts the real client IP from CapturedData, falling back to r.RemoteAddr.
func (mw *Middleware) resolveClientIP(r *http.Request) netip.Addr {
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
@@ -481,6 +508,9 @@ func (mw *Middleware) handleOAuthCallbackError(w http.ResponseWriter, r *http.Re
// forwardWithSessionCookie checks for a valid session cookie and, if found,
// sets the user identity on the request context and forwards to the next handler.
// A signature-valid cookie is not on its own a grant: an OIDC session must also
// carry a group the service allows, so a token cannot be replayed past the
// group check that gated the login it came from.
func (mw *Middleware) forwardWithSessionCookie(w http.ResponseWriter, r *http.Request, host string, config DomainConfig, next http.Handler) bool {
cookie, err := r.Cookie(auth.SessionCookieName)
if err != nil {
@@ -500,6 +530,14 @@ func (mw *Middleware) forwardWithSessionCookie(w http.ResponseWriter, r *http.Re
return false
}
if !sessionGroupsAllowed(config.AllowedGroups, auth.Method(method), groups) {
mw.logger.WithFields(log.Fields{
"host": host,
"user_id": userID,
}).Debug("session cookie rejected: groups claim does not intersect the service's allowed groups")
return false
}
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
cd.SetUserID(userID)
cd.SetUserEmail(email)
@@ -672,13 +710,9 @@ func (mw *Middleware) authenticateWithSchemes(w http.ResponseWriter, r *http.Req
var attemptedMethod string
for _, scheme := range config.Schemes {
token, promptData, err := scheme.Authenticate(r)
token, promptData, err := mw.authenticateScheme(r, config, scheme)
if err != nil {
mw.logger.WithField("scheme", scheme.Type().String()).Warnf("authentication infrastructure error: %v", err)
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
cd.SetOrigin(proxy.OriginAuth)
}
http.Error(w, "authentication service unavailable", http.StatusBadGateway)
mw.writeAuthenticationError(w, r, scheme.Type(), err)
return
}
@@ -779,9 +813,9 @@ func setSessionCookie(w http.ResponseWriter, token string, expiration time.Durat
func wasCredentialSubmitted(r *http.Request, method auth.Method) bool {
switch method {
case auth.MethodPIN:
return r.FormValue("pin") != ""
return credentialFormValue(r, pinFormId) != ""
case auth.MethodPassword:
return r.FormValue("password") != ""
return credentialFormValue(r, passwordFormId) != ""
case auth.MethodOIDC:
return r.URL.Query().Get(sessionTokenParam) != ""
}
@@ -802,6 +836,9 @@ type DomainSettings struct {
// of the schemes list.
Private bool
AppSecMode restrict.AppSecMode
// AllowedGroups restricts OIDC sessions to the given group ids; empty means
// unrestricted.
AllowedGroups []string
}
// AddDomain registers authentication schemes for the given domain. With schemes
@@ -814,6 +851,7 @@ func (mw *Middleware) AddDomain(domain string, settings DomainSettings) error {
IPRestrictions: settings.IPRestrictions,
Private: settings.Private,
AppSecMode: settings.AppSecMode,
AllowedGroups: groupSet(settings.AllowedGroups),
redactBodyFields: credentialFields,
redactHeaders: credentialHeaders(settings.Schemes),
// A credential can arrive in the query too: r.FormValue merges the URL
@@ -888,6 +926,50 @@ func (mw *Middleware) validateSessionToken(ctx context.Context, host, token stri
return &validationResult{UserID: userID, UserEmail: email, Valid: true, Groups: groups, GroupNames: groupNames}, nil
}
// groupSet builds the lookup set the cookie path consults, returning nil for an
// empty list so callers can test membership restriction with len().
func groupSet(groups []string) map[string]struct{} {
if len(groups) == 0 {
return nil
}
set := make(map[string]struct{}, len(groups))
for _, g := range groups {
if g != "" {
set[g] = struct{}{}
}
}
if len(set) == 0 {
return nil
}
return set
}
// sessionGroupsAllowed reports whether a session token's groups claim satisfies
// the service's allowed groups. Only OIDC sessions are gated: password, PIN and
// header credentials carry no group identity and are authorised by the secret
// itself, which mirrors how management validates them. A token minted before the
// groups claim existed carries none and is therefore denied on a group-restricted
// service, which sends the user back through login for a fresh decision. A method
// this build doesn't know carries no such argument, so it is denied.
func sessionGroupsAllowed(allowed map[string]struct{}, method auth.Method, groups []string) bool {
if len(allowed) == 0 {
return true
}
switch method {
case auth.MethodPassword, auth.MethodPIN, auth.MethodHeader:
return true
case auth.MethodOIDC:
for _, g := range groups {
if _, ok := allowed[g]; ok {
return true
}
}
return false
default:
return false
}
}
// stripSessionTokenParam returns the request URI with the session_token query
// parameter removed so it doesn't linger in the browser's address bar or history.
func stripSessionTokenParam(u *url.URL) string {
+1 -1
View File
@@ -35,7 +35,7 @@ func (Password) Type() auth.Method {
// so that it can be injected into a request from the UI so that
// authentication may be successful.
func (p Password) Authenticate(r *http.Request) (string, string, error) {
password := r.FormValue(passwordFormId)
password := credentialFormValue(r, passwordFormId)
if password == "" {
// No password submitted; return the form ID so the UI can prompt the user.
+1 -1
View File
@@ -35,7 +35,7 @@ func (Pin) Type() auth.Method {
// so that it can be injected into a request from the UI so that
// authentication may be successful.
func (p Pin) Authenticate(r *http.Request) (string, string, error) {
pin := r.FormValue(pinFormId)
pin := credentialFormValue(r, pinFormId)
if pin == "" {
// No PIN submitted; return the form ID so the UI can prompt the user.
+272
View File
@@ -0,0 +1,272 @@
package auth
import (
"context"
"errors"
"io"
"net/http"
"net/http/httptest"
"net/http/httptrace"
"net/netip"
"sync"
"testing"
"time"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
"github.com/netbirdio/netbird/proxy/internal/proxy"
"github.com/netbirdio/netbird/proxy/internal/restrict"
"github.com/netbirdio/netbird/shared/management/proto"
)
// switchableTunnelValidator flips the ValidateTunnelPeer verdict between requests.
type switchableTunnelValidator struct {
mu sync.Mutex
valid bool
}
func (s *switchableTunnelValidator) setValid(v bool) {
s.mu.Lock()
defer s.mu.Unlock()
s.valid = v
}
func (s *switchableTunnelValidator) ValidateSession(context.Context, *proto.ValidateSessionRequest, ...grpc.CallOption) (*proto.ValidateSessionResponse, error) {
return nil, errors.New("not used in this test")
}
func (s *switchableTunnelValidator) ValidateTunnelPeer(context.Context, *proto.ValidateTunnelPeerRequest, ...grpc.CallOption) (*proto.ValidateTunnelPeerResponse, error) {
s.mu.Lock()
defer s.mu.Unlock()
if !s.valid {
return &proto.ValidateTunnelPeerResponse{Valid: false, DeniedReason: "not_in_group"}, nil
}
return &proto.ValidateTunnelPeerResponse{
Valid: true,
UserId: "user-1",
SessionToken: "tunnel-session-token",
}, nil
}
// testServerHost is the domain key Protect derives from the httptest listener.
const testServerHost = "127.0.0.1"
var testTunnelIP = netip.MustParseAddr("100.90.1.14")
// startProtectedServer serves mw.Protect and stamps requests as overlay traffic.
func startProtectedServer(t *testing.T, mw *Middleware, clientIP netip.Addr, lookup TunnelLookupFunc, h2 bool) *httptest.Server {
t.Helper()
protected := mw.Protect(newPassthroughHandler())
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
cd := proxy.NewCapturedData("")
cd.SetClientIP(clientIP)
ctx := proxy.WithCapturedData(r.Context(), cd)
ctx = WithTunnelLookup(ctx, lookup)
protected.ServeHTTP(w, r.WithContext(ctx))
})
srv := httptest.NewUnstartedServer(handler)
if h2 {
srv.EnableHTTP2 = true
srv.StartTLS()
} else {
srv.Start()
}
t.Cleanup(srv.Close)
return srv
}
// tracedResponse is what a test observes from one client round trip.
type tracedResponse struct {
status int
protoMajor int
close bool
connection string
cacheControl string
reused bool
}
// doTraced GETs url and reports whether the connection that served it was reused.
func doTraced(t *testing.T, client *http.Client, url string) tracedResponse {
t.Helper()
var reused bool
trace := &httptrace.ClientTrace{
GotConn: func(info httptrace.GotConnInfo) { reused = info.Reused },
}
req, err := http.NewRequestWithContext(httptrace.WithClientTrace(context.Background(), trace), http.MethodGet, url, nil)
require.NoError(t, err)
resp, err := client.Do(req)
require.NoError(t, err)
defer func() { require.NoError(t, resp.Body.Close()) }()
_, err = io.Copy(io.Discard, resp.Body)
require.NoError(t, err)
return tracedResponse{
status: resp.StatusCode,
protoMajor: resp.ProtoMajor,
close: resp.Close,
connection: resp.Header.Get("Connection"),
cacheControl: resp.Header.Get("Cache-Control"),
reused: reused,
}
}
func acceptAllLookup(_ netip.Addr) (PeerIdentity, bool) {
return PeerIdentity{TunnelIP: testTunnelIP}, true
}
func newPrivateMiddleware(t *testing.T, validator SessionValidator, ipRestrictions *restrict.Filter) *Middleware {
t.Helper()
mw := NewMiddleware(log.StandardLogger(), validator, nil)
kp := generateTestKeyPair(t)
require.NoError(t, mw.AddDomain(testServerHost, DomainSettings{SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acct-1", ServiceID: "svc-1", IPRestrictions: ipRestrictions, Private: true}))
return mw
}
// A rejected tunnel peer must emit the exact lowercase "close" token h2 matches on.
func TestProtect_PrivateService_DeniedSetsCloseHeaders(t *testing.T) {
mw := newPrivateMiddleware(t, &switchableTunnelValidator{}, nil)
handler := mw.Protect(newPassthroughHandler())
cd := proxy.NewCapturedData("")
cd.SetClientIP(testTunnelIP)
req := httptest.NewRequest(http.MethodGet, "http://"+testServerHost+"/", nil)
req.RemoteAddr = testTunnelIP.String() + ":5000"
req = req.WithContext(WithTunnelLookup(proxy.WithCapturedData(req.Context(), cd), acceptAllLookup))
rec := httptest.NewRecorder()
handler.ServeHTTP(rec, req)
assert.Equal(t, http.StatusForbidden, rec.Code)
assert.Equal(t, "close", rec.Header().Get("Connection"), "private denial must ask the client to drop the connection")
assert.Equal(t, "no-store", rec.Header().Get("Cache-Control"), "private denial must not be cacheable")
}
// A denied client must not keep reusing the warm socket after joining the overlay.
func TestPrivateDeny_HTTP1_ClosesConnection(t *testing.T) {
validator := &switchableTunnelValidator{}
mw := newPrivateMiddleware(t, validator, nil)
srv := startProtectedServer(t, mw, testTunnelIP, acceptAllLookup, false)
client := srv.Client()
resp := doTraced(t, client, srv.URL)
assert.Equal(t, http.StatusForbidden, resp.status)
assert.Equal(t, 1, resp.protoMajor, "plain httptest server must speak HTTP/1.1")
// The Go client folds "Connection: close" into resp.close and drops the header.
assert.True(t, resp.close, "private denial must make the client mark the connection as not reusable")
assert.Equal(t, "no-store", resp.cacheControl, "private denial must not be cacheable")
validator.setValid(true)
resp2 := doTraced(t, client, srv.URL)
assert.Equal(t, http.StatusOK, resp2.status, "the retry must reach the upstream once the peer is valid")
assert.False(t, resp2.reused, "the retry must open a new connection")
}
// On HTTP/2 the header becomes a GOAWAY and the retry must use a new connection.
func TestPrivateDeny_HTTP2_SendsGoAway(t *testing.T) {
validator := &switchableTunnelValidator{}
mw := newPrivateMiddleware(t, validator, nil)
srv := startProtectedServer(t, mw, testTunnelIP, acceptAllLookup, true)
client := srv.Client()
resp := doTraced(t, client, srv.URL)
require.Equal(t, 2, resp.protoMajor, "test client must negotiate HTTP/2")
assert.Equal(t, http.StatusForbidden, resp.status)
assert.Empty(t, resp.connection, "HTTP/2 must not carry a Connection header on the wire")
assert.Equal(t, "no-store", resp.cacheControl)
validator.setValid(true)
resp2 := doTraced(t, client, srv.URL)
assert.Equal(t, 2, resp2.protoMajor)
assert.Equal(t, http.StatusOK, resp2.status, "the retry must reach the upstream once the peer is valid")
assert.False(t, resp2.reused, "GOAWAY must retire the connection so the retry opens a new one")
}
// Legitimate private traffic keeps its keep-alive connection.
func TestPrivateAllow_KeepsConnection(t *testing.T) {
validator := &switchableTunnelValidator{valid: true}
mw := newPrivateMiddleware(t, validator, nil)
srv := startProtectedServer(t, mw, testTunnelIP, acceptAllLookup, false)
client := srv.Client()
resp := doTraced(t, client, srv.URL)
assert.Equal(t, http.StatusOK, resp.status)
assert.Empty(t, resp.connection, "an allowed private request must not close the connection")
resp2 := doTraced(t, client, srv.URL)
assert.Equal(t, http.StatusOK, resp2.status)
assert.True(t, resp2.reused, "allowed private traffic must keep reusing the connection")
}
// Public denials keep the connection open; only private services change.
func TestPublicDeny_KeepsConnection(t *testing.T) {
mw := NewMiddleware(log.StandardLogger(), nil, nil)
filter := restrict.ParseFilter(restrict.FilterConfig{AllowedCIDRs: []string{"10.0.0.0/8"}})
require.NoError(t, mw.AddDomain(testServerHost, DomainSettings{AccountID: "acct-1", ServiceID: "svc-1", IPRestrictions: filter}))
srv := startProtectedServer(t, mw, netip.MustParseAddr("192.168.1.1"), nil, false)
client := srv.Client()
resp := doTraced(t, client, srv.URL)
assert.Equal(t, http.StatusForbidden, resp.status)
assert.Empty(t, resp.connection, "public denial must not close the connection")
assert.Empty(t, resp.cacheControl, "public denial must not gain cache headers")
resp2 := doTraced(t, client, srv.URL)
assert.Equal(t, http.StatusForbidden, resp2.status)
assert.True(t, resp2.reused, "public denials must keep reusing the connection")
}
// IP restriction denials on a private service must close the connection too.
func TestCheckIPRestrictions_PrivateDenialClosesConnection(t *testing.T) {
filter := restrict.ParseFilter(restrict.FilterConfig{AllowedCIDRs: []string{"100.64.0.0/16"}})
mw := newPrivateMiddleware(t, &switchableTunnelValidator{valid: true}, filter)
handler := mw.Protect(newPassthroughHandler())
tests := []struct {
name string
remoteAddr string
}{
{"denied by CIDR", "100.65.5.6:5000"},
{"unresolvable client address", "not-an-ip:1234"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "http://"+testServerHost+"/", nil)
req.RemoteAddr = tt.remoteAddr
rec := httptest.NewRecorder()
handler.ServeHTTP(rec, req)
assert.Equal(t, http.StatusForbidden, rec.Code)
assert.Equal(t, "close", rec.Header().Get("Connection"), "private IP-restriction denial must close the connection")
assert.Equal(t, "no-store", rec.Header().Get("Cache-Control"), "private IP-restriction denial must not be cacheable")
})
}
}
func TestCheckIPRestrictions_PublicDenialKeepsConnection(t *testing.T) {
mw := NewMiddleware(log.StandardLogger(), nil, nil)
filter := restrict.ParseFilter(restrict.FilterConfig{AllowedCIDRs: []string{"10.0.0.0/8"}})
require.NoError(t, mw.AddDomain(testServerHost, DomainSettings{AccountID: "acct-1", ServiceID: "svc-1", IPRestrictions: filter}))
handler := mw.Protect(newPassthroughHandler())
tests := []struct {
name string
remoteAddr string
}{
{"denied by CIDR", "192.168.1.1:5000"},
{"unresolvable client address", "not-an-ip:1234"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
req := httptest.NewRequest(http.MethodGet, "http://"+testServerHost+"/", nil)
req.RemoteAddr = tt.remoteAddr
rec := httptest.NewRecorder()
handler.ServeHTTP(rec, req)
assert.Equal(t, http.StatusForbidden, rec.Code)
assert.Empty(t, rec.Header().Get("Connection"), "public IP-restriction denial must not close the connection")
assert.Empty(t, rec.Header().Get("Cache-Control"), "public IP-restriction denial must not gain cache headers")
})
}
}
+167
View File
@@ -0,0 +1,167 @@
package auth
import (
"context"
"crypto/tls"
"net/http"
"net/http/httptest"
"testing"
"time"
log "github.com/sirupsen/logrus"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/grpc"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
"github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/proxy/internal/proxy"
"github.com/netbirdio/netbird/shared/management/proto"
)
// denyingSessionValidator mimics management for a user who completed OIDC login
// but is outside the service's distribution groups: ValidateSession denies.
type denyingSessionValidator struct {
calls int
}
func (d *denyingSessionValidator) ValidateSession(context.Context, *proto.ValidateSessionRequest, ...grpc.CallOption) (*proto.ValidateSessionResponse, error) {
d.calls++
return &proto.ValidateSessionResponse{Valid: false, UserId: "user-1", DeniedReason: "not_in_group"}, nil
}
func (d *denyingSessionValidator) ValidateTunnelPeer(context.Context, *proto.ValidateTunnelPeerRequest, ...grpc.CallOption) (*proto.ValidateTunnelPeerResponse, error) {
return &proto.ValidateTunnelPeerResponse{Valid: false}, nil
}
// TestProtect_SelfInstalledCookieCannotBypassGroupCheck is the regression guard
// for the group-authorisation bypass: a user denied at login still holds the raw
// session token from the ?session_token= redirect, so pasting it into the
// nb_session cookie must not buy access. The cookie path validated only the JWT
// signature, which turned the token management had already refused into a bearer
// credential for the service.
func TestProtect_SelfInstalledCookieCannotBypassGroupCheck(t *testing.T) {
validator := &denyingSessionValidator{}
mw := NewMiddleware(log.StandardLogger(), validator, nil)
kp := generateTestKeyPair(t)
oidc := &stubScheme{method: auth.MethodOIDC, authFn: func(r *http.Request) (string, string, error) {
return r.URL.Query().Get("session_token"), "https://idp.example/authorize", nil
}}
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{oidc}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acct-1", ServiceID: "svc-1", AllowedGroups: []string{"grp-allowed"}}))
// The token a denied user gets to see: validly signed for this service and
// domain, but carrying no group the service allows.
token, err := sessionkey.SignToken(kp.PrivateKey, "user-1", "john.doe@example.com", "example.com", auth.MethodOIDC, nil, nil, time.Hour)
require.NoError(t, err)
backendHits := 0
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
backendHits++
w.WriteHeader(http.StatusOK)
}))
t.Run("token in the callback URL is denied", func(t *testing.T) {
rec := serveWithCookie(t, handler, "https://example.com/?session_token="+token, nil)
assert.Equal(t, http.StatusForbidden, rec.Code, "group check must deny the login")
assert.Empty(t, rec.Result().Cookies(), "a denied login must not install a session cookie")
})
t.Run("same token pasted into the session cookie is denied", func(t *testing.T) {
rec := serveWithCookie(t, handler, "https://example.com/", &http.Cookie{Name: auth.SessionCookieName, Value: token})
assert.NotEqual(t, http.StatusOK, rec.Code, "a self-installed cookie must not reach the backend")
assert.Equal(t, 0, backendHits, "backend must never be reached without an allowed group")
})
}
// TestProtect_SessionCookieWithAllowedGroupPassesThrough is the positive half of
// the group gate: a member of an allowed group keeps the cookie fast-path, with
// no management round-trip.
func TestProtect_SessionCookieWithAllowedGroupPassesThrough(t *testing.T) {
validator := &denyingSessionValidator{}
mw := NewMiddleware(log.StandardLogger(), validator, nil)
kp := generateTestKeyPair(t)
oidc := &stubScheme{method: auth.MethodOIDC}
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{oidc}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acct-1", ServiceID: "svc-1", AllowedGroups: []string{"grp-other", "grp-allowed"}}))
token, err := sessionkey.SignToken(kp.PrivateKey, "user-2", "jane@example.com", "example.com", auth.MethodOIDC,
[]string{"grp-unrelated", "grp-allowed"}, []string{"Unrelated", "Allowed"}, time.Hour)
require.NoError(t, err)
handler := mw.Protect(newPassthroughHandler())
rec := serveWithCookie(t, handler, "https://example.com/", &http.Cookie{Name: auth.SessionCookieName, Value: token})
assert.Equal(t, http.StatusOK, rec.Code, "a cookie carrying an allowed group must pass through")
assert.Equal(t, 0, validator.calls, "the cookie fast-path must not call management")
}
// TestProtect_NonOIDCSessionCookieIgnoresGroupRestriction locks the scope of the
// gate: PIN, password and header credentials carry no group identity and are
// authorised by the secret itself, exactly as management validates them.
func TestProtect_NonOIDCSessionCookieIgnoresGroupRestriction(t *testing.T) {
mw := NewMiddleware(log.StandardLogger(), nil, nil)
kp := generateTestKeyPair(t)
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acct-1", ServiceID: "svc-1", AllowedGroups: []string{"grp-allowed"}}))
token, err := sessionkey.SignToken(kp.PrivateKey, "pin-user", "", "example.com", auth.MethodPIN, nil, nil, time.Hour)
require.NoError(t, err)
handler := mw.Protect(newPassthroughHandler())
rec := serveWithCookie(t, handler, "https://example.com/", &http.Cookie{Name: auth.SessionCookieName, Value: token})
assert.Equal(t, http.StatusOK, rec.Code, "a PIN session must not be gated on OIDC group membership")
}
func TestSessionGroupsAllowed(t *testing.T) {
allowed := groupSet([]string{"a", "b"})
tests := []struct {
name string
allowed map[string]struct{}
method auth.Method
groups []string
want bool
}{
{"unrestricted service allows a groupless token", nil, auth.MethodOIDC, nil, true},
{"restricted service allows an intersecting token", allowed, auth.MethodOIDC, []string{"c", "b"}, true},
{"restricted service denies a disjoint token", allowed, auth.MethodOIDC, []string{"c"}, false},
{"restricted service denies a groupless token", allowed, auth.MethodOIDC, nil, false},
{"restricted service ignores a pin token", allowed, auth.MethodPIN, nil, true},
{"restricted service ignores a password token", allowed, auth.MethodPassword, nil, true},
{"restricted service ignores a header token", allowed, auth.MethodHeader, nil, true},
{"restricted service denies an unknown method", allowed, auth.Method("totp"), []string{"a"}, false},
{"restricted service denies a token with no method", allowed, auth.Method(""), []string{"a"}, false},
}
for _, tc := range tests {
t.Run(tc.name, func(t *testing.T) {
assert.Equal(t, tc.want, sessionGroupsAllowed(tc.allowed, tc.method, tc.groups))
})
}
}
func TestGroupSetDropsEmptyEntries(t *testing.T) {
assert.Nil(t, groupSet(nil), "no groups means unrestricted")
assert.Nil(t, groupSet([]string{"", ""}), "blank ids must not restrict access to nothing reachable")
assert.Equal(t, map[string]struct{}{"a": {}}, groupSet([]string{"a", ""}))
}
// serveWithCookie drives the middleware over TLS with captured data attached,
// optionally carrying a session cookie.
func serveWithCookie(t *testing.T, handler http.Handler, url string, cookie *http.Cookie) *httptest.ResponseRecorder {
t.Helper()
req := httptest.NewRequest(http.MethodGet, url, nil)
req.TLS = &tls.ConnectionState{}
if cookie != nil {
req.AddCookie(cookie)
}
req = req.WithContext(proxy.WithCapturedData(req.Context(), proxy.NewCapturedData("")))
rec := httptest.NewRecorder()
handler.ServeHTTP(rec, req)
return rec
}