mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-08 14:39:09 +02:00
[management,proxy] Use single-use codes for OIDC session handoff (#7635)
* Generalize PKCE verifier store into SingleUseStore * Generalize PKCE verifier store into SingleUseStore * Extend single-use store to generate one-time retrieval codes * Hand off proxy OIDC session via one-time code instead of URL token * Use the single-use store in integration tests * Read active proxy versions by cluster * Detect proxy clusters that support session codes * Bind OIDC session handoff mode to signed state * Deprecate legacy OIDC session token handoff * Remove unrelated session code test stub * fix tests * fix merge * Fix session code compatibility detection * Isolate proxy session codes in shared cache * bump min session version
This commit is contained in:
@@ -20,6 +20,7 @@ type Manager interface {
|
||||
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool
|
||||
CleanupStale(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error)
|
||||
CountAccountProxies(ctx context.Context, accountID string) (int64, error)
|
||||
|
||||
@@ -8,6 +8,7 @@ import (
|
||||
"go.opentelemetry.io/otel/metric"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
nbversion "github.com/netbirdio/netbird/version"
|
||||
)
|
||||
|
||||
// store defines the interface for proxy persistence operations
|
||||
@@ -22,6 +23,7 @@ type store interface {
|
||||
GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error)
|
||||
CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error)
|
||||
CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error)
|
||||
@@ -29,6 +31,8 @@ type store interface {
|
||||
DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error
|
||||
}
|
||||
|
||||
const minSessionCodeVersion = "0.81.0"
|
||||
|
||||
// Manager handles all proxy operations
|
||||
type Manager struct {
|
||||
store store
|
||||
@@ -145,6 +149,22 @@ func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string)
|
||||
return m.store.GetClusterSupportsPrivate(ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterSupportsSessionCode reports whether all active proxies support session codes.
|
||||
func (m Manager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
|
||||
versions, err := m.store.GetActiveProxyVersions(ctx, clusterAddr)
|
||||
if err != nil || len(versions) == 0 {
|
||||
return false
|
||||
}
|
||||
|
||||
for _, version := range versions {
|
||||
if supported, err := nbversion.MeetsMinVersion(minSessionCodeVersion, version); err != nil || !supported {
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// CleanupStale removes proxies that haven't sent heartbeat in the specified duration
|
||||
func (m *Manager) CleanupStale(ctx context.Context, inactivityDuration time.Duration) error {
|
||||
if err := m.store.CleanupStaleProxies(ctx, inactivityDuration); err != nil {
|
||||
|
||||
@@ -22,6 +22,7 @@ type mockStore struct {
|
||||
updateProxyHeartbeatFunc func(ctx context.Context, p *proxy.Proxy) error
|
||||
getActiveProxyClusterAddressesFunc func(ctx context.Context) ([]string, error)
|
||||
getActiveProxyClusterAddressesForAccFunc func(ctx context.Context, accountID string) ([]string, error)
|
||||
getActiveProxyVersionsFunc func(ctx context.Context, clusterAddress string) ([]string, error)
|
||||
cleanupStaleProxiesFunc func(ctx context.Context, d time.Duration) error
|
||||
getProxyByAccountIDFunc func(ctx context.Context, accountID string) (*proxy.Proxy, error)
|
||||
countProxiesByAccountIDFunc func(ctx context.Context, accountID string) (int64, error)
|
||||
@@ -104,6 +105,12 @@ func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *boo
|
||||
func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
func (m *mockStore) GetActiveProxyVersions(ctx context.Context, clusterAddress string) ([]string, error) {
|
||||
if m.getActiveProxyVersionsFunc != nil {
|
||||
return m.getActiveProxyVersionsFunc(ctx, clusterAddress)
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func newTestManager(s store) *Manager {
|
||||
meter := noop.NewMeterProvider().Meter("test")
|
||||
@@ -114,6 +121,34 @@ func newTestManager(s store) *Manager {
|
||||
return m
|
||||
}
|
||||
|
||||
func TestClusterSupportsSessionCode(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
versions []string
|
||||
storeErr error
|
||||
want bool
|
||||
}{
|
||||
{name: "all supported", versions: []string{"0.81.0", "0.81.2"}, want: true},
|
||||
{name: "one old proxy", versions: []string{"0.81.0", "0.80.0"}},
|
||||
{name: "missing version", versions: []string{"0.81.0", ""}},
|
||||
{name: "no active proxies"},
|
||||
{name: "store error", storeErr: errors.New("db error")},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
s := &mockStore{
|
||||
getActiveProxyVersionsFunc: func(_ context.Context, _ string) ([]string, error) {
|
||||
return tt.versions, tt.storeErr
|
||||
},
|
||||
}
|
||||
|
||||
got := newTestManager(s).ClusterSupportsSessionCode(context.Background(), "cluster.example.com")
|
||||
assert.Equal(t, tt.want, got)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestConnect_WithAccountID(t *testing.T) {
|
||||
accountID := "acc-123"
|
||||
|
||||
|
||||
@@ -112,6 +112,20 @@ func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterSupportsSessionCode mocks base method.
|
||||
func (m *MockManager) ClusterSupportsSessionCode(ctx context.Context, clusterAddr string) bool {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ClusterSupportsSessionCode", ctx, clusterAddr)
|
||||
ret0, _ := ret[0].(bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// ClusterSupportsSessionCode indicates an expected call of ClusterSupportsSessionCode.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsSessionCode(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsSessionCode", reflect.TypeOf((*MockManager)(nil).ClusterSupportsSessionCode), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// Connect mocks base method.
|
||||
func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAddress, ipAddress, version string, accountID *string, capabilities *Capabilities) (*Proxy, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
@@ -27,8 +27,6 @@ import (
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/peers"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
@@ -42,6 +40,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/users"
|
||||
proxyauth "github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/shared/hash/argon2id"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
nbstatus "github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
@@ -142,7 +141,7 @@ type ProxyServiceServer struct {
|
||||
// OIDC configuration for proxy authentication
|
||||
oidcConfig ProxyOIDCConfig
|
||||
|
||||
// Store for PKCE verifiers
|
||||
// singleUseStore backs both PKCE verifiers and OIDC session exchange codes.
|
||||
singleUseStore *SingleUseStore
|
||||
|
||||
// tokenTTL is the lifetime of one-time tokens generated for proxy
|
||||
@@ -158,6 +157,13 @@ type ProxyServiceServer struct {
|
||||
|
||||
const pkceVerifierTTL = 10 * time.Minute
|
||||
|
||||
const sessionCodeTTL = 60 * time.Second
|
||||
|
||||
const sessionCodeCacheNamespace = "proxy:session"
|
||||
|
||||
// The signed nonce binds the handoff mode without changing the state format.
|
||||
const sessionCodeNoncePrefix = "code."
|
||||
|
||||
const defaultProxyTokenTTL = 5 * time.Minute
|
||||
|
||||
const defaultSnapshotBatchSize = 500
|
||||
@@ -306,6 +312,16 @@ func (s *ProxyServiceServer) proxyConnectAuthorizer() ProxyConnectAuthorizer {
|
||||
return s.connectAuthorizer
|
||||
}
|
||||
|
||||
// GenerateSessionCode creates a single-use code for the given session token.
|
||||
func (s *ProxyServiceServer) GenerateSessionCode(sessionToken string) (code string, ok bool) {
|
||||
code, err := s.singleUseStore.Generate(sessionCodeCacheNamespace, sessionToken, sessionCodeTTL)
|
||||
if err != nil {
|
||||
log.WithError(err).Error("failed to generate proxy session code")
|
||||
return "", false
|
||||
}
|
||||
return code, true
|
||||
}
|
||||
|
||||
// CheckLLMPolicyLimits is the pre-flight policy gate the proxy calls before
|
||||
// forwarding an LLM request upstream. Delegates to the agent-network selector,
|
||||
// which scores applicable policies by remaining headroom and returns the
|
||||
@@ -1536,18 +1552,20 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
|
||||
log.WithContext(ctx).Errorf("failed to get account services: %v", err)
|
||||
return nil, status.Errorf(codes.FailedPrecondition, "get account services: %v", err)
|
||||
}
|
||||
var found bool
|
||||
var matchedService *rpservice.Service
|
||||
for _, service := range services {
|
||||
if service.Domain == redirectURL.Hostname() {
|
||||
found = true
|
||||
matchedService = service
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
if matchedService == nil {
|
||||
log.WithContext(ctx).Debugf("OIDC redirect URL %q does not match any service domain", redirectURL.Hostname())
|
||||
return nil, status.Errorf(codes.FailedPrecondition, "service not found in store")
|
||||
}
|
||||
|
||||
useSessionCode := s.proxyManager.ClusterSupportsSessionCode(ctx, matchedService.ProxyCluster)
|
||||
|
||||
provider, err := oidc.NewProvider(ctx, s.oidcConfig.Issuer)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to create OIDC provider: %v", err)
|
||||
@@ -1567,9 +1585,12 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
|
||||
return nil, status.Errorf(codes.Internal, "generate nonce: %v", err)
|
||||
}
|
||||
nonceB64 := base64.URLEncoding.EncodeToString(nonce)
|
||||
if useSessionCode {
|
||||
nonceB64 = sessionCodeNoncePrefix + nonceB64
|
||||
}
|
||||
|
||||
// Using an HMAC here to avoid redirection state being modified.
|
||||
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
|
||||
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
|
||||
payload := redirectURL.String() + "|" + nonceB64
|
||||
hmacSum := s.generateHMAC(payload)
|
||||
state := fmt.Sprintf("%s|%s|%s", base64.URLEncoding.EncodeToString([]byte(redirectURL.String())), nonceB64, hmacSum)
|
||||
@@ -1612,15 +1633,12 @@ func (s *ProxyServiceServer) generateHMAC(input string) string {
|
||||
return hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
// ValidateState validates the state parameter from an OAuth callback.
|
||||
// Returns the original redirect URL if valid, or an error if invalid.
|
||||
// The HMAC is verified before consuming the PKCE verifier to prevent
|
||||
// an attacker from invalidating a legitimate user's auth flow.
|
||||
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, err error) {
|
||||
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
|
||||
// ValidateState validates and consumes an OIDC state.
|
||||
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, useSessionCode bool, err error) {
|
||||
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
|
||||
parts := strings.Split(state, "|")
|
||||
if len(parts) != 3 {
|
||||
return "", "", errors.New("invalid state format")
|
||||
return "", "", false, errors.New("invalid state format")
|
||||
}
|
||||
|
||||
encodedURL := parts[0]
|
||||
@@ -1629,7 +1647,7 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
|
||||
|
||||
redirectURLBytes, err := base64.URLEncoding.DecodeString(encodedURL)
|
||||
if err != nil {
|
||||
return "", "", fmt.Errorf("invalid state encoding: %w", err)
|
||||
return "", "", false, fmt.Errorf("invalid state encoding: %w", err)
|
||||
}
|
||||
redirectURL = string(redirectURLBytes)
|
||||
|
||||
@@ -1637,16 +1655,17 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
|
||||
expectedHMAC := s.generateHMAC(payload)
|
||||
|
||||
if !hmac.Equal([]byte(providedHMAC), []byte(expectedHMAC)) {
|
||||
return "", "", errors.New("invalid state signature")
|
||||
return "", "", false, errors.New("invalid state signature")
|
||||
}
|
||||
useSessionCode = strings.HasPrefix(nonce, sessionCodeNoncePrefix)
|
||||
|
||||
// Consume the PKCE verifier only after HMAC validation passes.
|
||||
verifier, ok := s.singleUseStore.LoadAndDelete(state)
|
||||
if !ok {
|
||||
return "", "", errors.New("no verifier for state")
|
||||
return "", "", false, errors.New("no verifier for state")
|
||||
}
|
||||
|
||||
return verifier, redirectURL, nil
|
||||
return verifier, redirectURL, useSessionCode, nil
|
||||
}
|
||||
|
||||
// Denied reasons reported to the proxy when access is refused because of the
|
||||
@@ -1851,7 +1870,20 @@ func (s *ProxyServiceServer) getAccountServiceByDomain(ctx context.Context, acco
|
||||
// ValidateSession validates a session token and checks if the user has access to the domain.
|
||||
func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.ValidateSessionRequest) (*proto.ValidateSessionResponse, error) {
|
||||
domain := req.GetDomain()
|
||||
sessionToken := req.GetSessionToken()
|
||||
sessionToken := req.GetSessionToken() //nolint:staticcheck
|
||||
|
||||
// A one-time code from the OIDC callback is redeemed here for the durable
|
||||
// token, so the token never travels in a redirect URL. The redeemed token
|
||||
// is returned to the proxy (mintedToken) to install as the session cookie.
|
||||
mintedToken := ""
|
||||
if code := req.GetSessionCode(); code != "" {
|
||||
redeemed, found := s.singleUseStore.LoadAndDelete(singleUseCacheKey(sessionCodeCacheNamespace, code))
|
||||
if !found {
|
||||
return deniedSessionResponse("invalid or expired session code"), nil
|
||||
}
|
||||
sessionToken = redeemed
|
||||
mintedToken = redeemed
|
||||
}
|
||||
|
||||
if domain == "" || sessionToken == "" {
|
||||
return deniedSessionResponse("missing domain or session_token"), nil
|
||||
@@ -1924,6 +1956,7 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val
|
||||
UserEmail: user.Email,
|
||||
PeerGroupIds: groupIDs,
|
||||
PeerGroupNames: groupNames,
|
||||
SessionToken: mintedToken,
|
||||
}, nil
|
||||
}
|
||||
|
||||
|
||||
@@ -313,7 +313,7 @@ func TestValidateState_RejectsOldTwoPartFormat(t *testing.T) {
|
||||
err := s.singleUseStore.Store("base64url|hmac", "test", 10*time.Minute)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, _, err = s.ValidateState("base64url|hmac")
|
||||
_, _, _, err = s.ValidateState("base64url|hmac")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid state format")
|
||||
}
|
||||
@@ -385,11 +385,35 @@ func TestValidateState_RejectsInvalidHMAC(t *testing.T) {
|
||||
err := s.singleUseStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute)
|
||||
require.NoError(t, err)
|
||||
|
||||
_, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
|
||||
_, _, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid state signature")
|
||||
}
|
||||
|
||||
func TestSessionCodeCannotConsumeOIDCState(t *testing.T) {
|
||||
const verifier = "pkce-verifier"
|
||||
|
||||
store := NewSingleUseStore(context.Background(), testCacheStore(t))
|
||||
server := &ProxyServiceServer{
|
||||
oidcConfig: ProxyOIDCConfig{
|
||||
HMACKey: []byte("test-hmac-key"),
|
||||
},
|
||||
singleUseStore: store,
|
||||
}
|
||||
state := generateState(server, "https://service.example.com/callback")
|
||||
require.NoError(t, store.Store(state, verifier, time.Minute))
|
||||
|
||||
response, err := server.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
|
||||
SessionCode: state,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.False(t, response.GetValid())
|
||||
|
||||
gotVerifier, _, _, err := server.ValidateState(state)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, verifier, gotVerifier)
|
||||
}
|
||||
|
||||
func TestSendServiceUpdateToCluster_FiltersOnCapability(t *testing.T) {
|
||||
tokenStore := NewOneTimeTokenStore(context.Background(), testCacheStore(t))
|
||||
|
||||
|
||||
@@ -2,6 +2,8 @@ package grpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
@@ -33,6 +35,24 @@ func (s *SingleUseStore) Store(key, value string, ttl time.Duration) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Generate stores a value under a namespaced random key and returns the random key.
|
||||
func (s *SingleUseStore) Generate(namespace, value string, ttl time.Duration) (string, error) {
|
||||
buf := make([]byte, 32)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", fmt.Errorf("generate single-use key: %w", err)
|
||||
}
|
||||
|
||||
key := base64.RawURLEncoding.EncodeToString(buf)
|
||||
if err := s.Store(singleUseCacheKey(namespace, key), value, ttl); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return key, nil
|
||||
}
|
||||
|
||||
func singleUseCacheKey(namespace, key string) string {
|
||||
return namespace + ":" + key
|
||||
}
|
||||
|
||||
// LoadAndDelete retrieves and removes the value for a key.
|
||||
func (s *SingleUseStore) LoadAndDelete(key string) (string, bool) {
|
||||
value, found, err := s.cache.GetDel(s.ctx, key)
|
||||
|
||||
@@ -83,3 +83,40 @@ func TestSingleUseStoreLoadAndDelete(t *testing.T) {
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestSingleUseStore_GenerateAndConsumeOnce(t *testing.T) {
|
||||
const namespace = "test"
|
||||
s := NewSingleUseStore(context.Background(), testCacheStore(t))
|
||||
|
||||
key, err := s.Generate(namespace, "the-value", time.Minute)
|
||||
if err != nil {
|
||||
t.Fatalf("generate: %v", err)
|
||||
}
|
||||
if key == "" || key == "the-value" {
|
||||
t.Fatalf("unexpected key %q", key)
|
||||
}
|
||||
|
||||
value, found := s.LoadAndDelete(singleUseCacheKey(namespace, key))
|
||||
if !found || value != "the-value" {
|
||||
t.Fatalf("expected to load the stored value, got %q found=%v", value, found)
|
||||
}
|
||||
|
||||
if _, found := s.LoadAndDelete(singleUseCacheKey(namespace, key)); found {
|
||||
t.Fatal("value must be consumed on first LoadAndDelete")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSingleUseStore_GenerateUniqueKeys(t *testing.T) {
|
||||
s := NewSingleUseStore(context.Background(), testCacheStore(t))
|
||||
a, err := s.Generate("test", "v", time.Minute)
|
||||
if err != nil {
|
||||
t.Fatalf("generate a: %v", err)
|
||||
}
|
||||
b, err := s.Generate("test", "v", time.Minute)
|
||||
if err != nil {
|
||||
t.Fatalf("generate b: %v", err)
|
||||
}
|
||||
if a == b {
|
||||
t.Fatal("generated keys must be distinct")
|
||||
}
|
||||
}
|
||||
|
||||
@@ -634,6 +634,10 @@ func (m *testValidateSessionProxyManager) ClusterSupportsPrivate(_ context.Conte
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *testValidateSessionProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
|
||||
return false
|
||||
}
|
||||
|
||||
type testValidateSessionUsersManager struct {
|
||||
store store.Store
|
||||
}
|
||||
@@ -662,3 +666,47 @@ func (m *testValidateSessionUsersManager) GetUserWithGroups(ctx context.Context,
|
||||
}
|
||||
return user, groups, nil
|
||||
}
|
||||
|
||||
func TestValidateSession_RedeemsSessionCode(t *testing.T) {
|
||||
setup := setupValidateSessionTest(t)
|
||||
defer setup.cleanup()
|
||||
|
||||
proxy, err := setup.store.GetServiceByID(context.Background(), store.LockingStrengthNone, "testAccountId", "testProxyId")
|
||||
require.NoError(t, err)
|
||||
|
||||
token := createSessionToken(t, proxy.SessionPrivateKey, "allowedUserId", "test-proxy.example.com")
|
||||
code, ok := setup.proxyService.GenerateSessionCode(token)
|
||||
require.True(t, ok)
|
||||
require.NotEqual(t, token, code, "code must not be the token itself")
|
||||
|
||||
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
|
||||
Domain: "test-proxy.example.com",
|
||||
SessionCode: code,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.True(t, resp.Valid, "redeemed code should authorize the user")
|
||||
assert.Equal(t, "allowedUserId", resp.UserId)
|
||||
assert.Equal(t, token, resp.GetSessionToken(), "response must carry the durable token for the cookie")
|
||||
|
||||
// Single-use: the same code must not redeem again.
|
||||
resp2, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
|
||||
Domain: "test-proxy.example.com",
|
||||
SessionCode: code,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.False(t, resp2.Valid, "a consumed code must be rejected")
|
||||
assert.Empty(t, resp2.GetSessionToken())
|
||||
}
|
||||
|
||||
func TestValidateSession_InvalidSessionCode(t *testing.T) {
|
||||
setup := setupValidateSessionTest(t)
|
||||
defer setup.cleanup()
|
||||
|
||||
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
|
||||
Domain: "test-proxy.example.com",
|
||||
SessionCode: "does-not-exist",
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.False(t, resp.Valid)
|
||||
assert.Empty(t, resp.GetSessionToken())
|
||||
}
|
||||
|
||||
@@ -59,7 +59,7 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
|
||||
|
||||
state := r.URL.Query().Get("state")
|
||||
|
||||
codeVerifier, originalURL, err := h.proxyService.ValidateState(state)
|
||||
codeVerifier, originalURL, useSessionCode, err := h.proxyService.ValidateState(state)
|
||||
if err != nil {
|
||||
log.WithError(err).Error("OAuth callback state validation failed")
|
||||
http.Error(w, "Invalid state parameter", http.StatusBadRequest)
|
||||
@@ -119,10 +119,19 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
|
||||
redirectURL.Scheme = "https"
|
||||
|
||||
query := redirectURL.Query()
|
||||
query.Set("session_token", sessionToken)
|
||||
if useSessionCode {
|
||||
code, ok := h.proxyService.GenerateSessionCode(sessionToken)
|
||||
if !ok {
|
||||
http.Error(w, "Failed to create session", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
query.Set("session_code", code)
|
||||
} else {
|
||||
query.Set("session_token", sessionToken)
|
||||
}
|
||||
redirectURL.RawQuery = query.Encode()
|
||||
|
||||
log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user with session token")
|
||||
log.WithField("redirect", redirectURL.Host).Debug("OAuth callback: redirecting user to proxy")
|
||||
http.Redirect(w, r, redirectURL.String(), http.StatusFound)
|
||||
}
|
||||
|
||||
|
||||
@@ -181,6 +181,10 @@ func (m *testAccessLogManager) GetAllAccessLogs(_ context.Context, _, _ string,
|
||||
}
|
||||
|
||||
func setupAuthCallbackTest(t *testing.T) *testSetup {
|
||||
return setupAuthCallbackTestWithProxyManager(t, testSessionCodeManager{})
|
||||
}
|
||||
|
||||
func setupAuthCallbackTestWithProxyManager(t *testing.T, proxyManager nbproxy.Manager) *testSetup {
|
||||
t.Helper()
|
||||
|
||||
ctx := context.Background()
|
||||
@@ -217,7 +221,7 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
|
||||
nil,
|
||||
usersManager,
|
||||
nil,
|
||||
nil,
|
||||
proxyManager,
|
||||
nil,
|
||||
)
|
||||
|
||||
@@ -242,6 +246,15 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
|
||||
}
|
||||
}
|
||||
|
||||
type testSessionCodeManager struct {
|
||||
nbproxy.Manager
|
||||
supported bool
|
||||
}
|
||||
|
||||
func (m testSessionCodeManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
|
||||
return m.supported
|
||||
}
|
||||
|
||||
func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store.Store) {
|
||||
t.Helper()
|
||||
|
||||
@@ -252,10 +265,11 @@ func createTestReverseProxies(t *testing.T, ctx context.Context, testStore store
|
||||
privKey := base64.StdEncoding.EncodeToString(priv)
|
||||
|
||||
testProxy := &service.Service{
|
||||
ID: "testProxyId",
|
||||
AccountID: "testAccountId",
|
||||
Name: "Test Proxy",
|
||||
Domain: "test-proxy.example.com",
|
||||
ID: "testProxyId",
|
||||
AccountID: "testAccountId",
|
||||
Name: "Test Proxy",
|
||||
Domain: "test-proxy.example.com",
|
||||
ProxyCluster: "cluster.example.com",
|
||||
Targets: []*service.Target{{
|
||||
Path: strPtr("/"),
|
||||
Host: "localhost",
|
||||
@@ -512,29 +526,56 @@ func createTestState(t *testing.T, ps *nbgrpc.ProxyServiceServer, redirectURL st
|
||||
}
|
||||
|
||||
func TestAuthCallback_UserAllowedToLogin(t *testing.T) {
|
||||
setup := setupAuthCallbackTest(t)
|
||||
defer setup.cleanup()
|
||||
tests := []struct {
|
||||
name string
|
||||
manager nbproxy.Manager
|
||||
wantParam string
|
||||
absentParam string
|
||||
}{
|
||||
{name: "legacy proxy", manager: testSessionCodeManager{}, wantParam: "session_token", absentParam: "session_code"},
|
||||
{name: "compatible proxy", manager: testSessionCodeManager{supported: true}, wantParam: "session_code", absentParam: "session_token"},
|
||||
}
|
||||
|
||||
setup.oidcServer.tokenSubject = "allowedUserId"
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
setup := setupAuthCallbackTestWithProxyManager(t, tt.manager)
|
||||
defer setup.cleanup()
|
||||
|
||||
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
|
||||
setup.oidcServer.tokenSubject = "allowedUserId"
|
||||
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
|
||||
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
|
||||
rec := httptest.NewRecorder()
|
||||
setup.router.ServeHTTP(rec, req)
|
||||
require.Equal(t, http.StatusFound, rec.Code)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
|
||||
rec := httptest.NewRecorder()
|
||||
location, err := url.Parse(rec.Header().Get("Location"))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, "test-proxy.example.com", location.Host)
|
||||
require.NotEmpty(t, location.Query().Get(tt.wantParam))
|
||||
require.Empty(t, location.Query().Get(tt.absentParam))
|
||||
require.Empty(t, location.Query().Get("error"))
|
||||
|
||||
setup.router.ServeHTTP(rec, req)
|
||||
if tt.wantParam == "session_code" {
|
||||
code := location.Query().Get("session_code")
|
||||
response, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
|
||||
Domain: location.Hostname(),
|
||||
SessionCode: code,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.True(t, response.GetValid())
|
||||
require.NotEmpty(t, response.GetSessionToken())
|
||||
require.NotEqual(t, code, response.GetSessionToken())
|
||||
|
||||
require.Equal(t, http.StatusFound, rec.Code)
|
||||
|
||||
location := rec.Header().Get("Location")
|
||||
require.NotEmpty(t, location)
|
||||
|
||||
parsedLocation, err := url.Parse(location)
|
||||
require.NoError(t, err)
|
||||
|
||||
require.Equal(t, "test-proxy.example.com", parsedLocation.Host)
|
||||
require.NotEmpty(t, parsedLocation.Query().Get("session_token"), "Should include session token")
|
||||
require.Empty(t, parsedLocation.Query().Get("error"), "Should not have error parameter")
|
||||
replayed, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
|
||||
Domain: location.Hostname(),
|
||||
SessionCode: code,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
require.False(t, replayed.GetValid())
|
||||
require.Empty(t, replayed.GetSessionToken())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestAuthCallback_UserDeniedByAccountStatus asserts that a user whose account
|
||||
|
||||
@@ -367,6 +367,21 @@ func (s *SqlStore) GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr s
|
||||
return s.getClusterUnanimousCapability(ctx, clusterAddr, "supports_crowdsec")
|
||||
}
|
||||
|
||||
// GetActiveProxyVersions returns every active proxy version in a cluster.
|
||||
func (s *SqlStore) GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) {
|
||||
var versions []string
|
||||
err := s.db.WithContext(ctx).
|
||||
Model(&proxy.Proxy{}).
|
||||
Where("LOWER(cluster_address) = LOWER(?) AND status = ? AND last_seen > ?",
|
||||
clusterAddr, proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)).
|
||||
Pluck("version", &versions).Error
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("failed to get active proxy versions for %s: %v", clusterAddr, err)
|
||||
return nil, status.Errorf(status.Internal, "get active proxy versions")
|
||||
}
|
||||
return versions, nil
|
||||
}
|
||||
|
||||
// getClusterUnanimousCapability returns an aggregated boolean capability
|
||||
// requiring all active proxies in the cluster to report true.
|
||||
func (s *SqlStore) getClusterUnanimousCapability(ctx context.Context, clusterAddr, column string) *bool {
|
||||
|
||||
@@ -337,6 +337,7 @@ type Store interface {
|
||||
GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error)
|
||||
CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetAllProxies(ctx context.Context) ([]*proxy.Proxy, error)
|
||||
DisconnectAllProxies(ctx context.Context) (int64, error)
|
||||
|
||||
@@ -1539,6 +1539,21 @@ func (mr *MockStoreMockRecorder) GetActiveProxyClusterAddressesForAccount(ctx, a
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveProxyClusterAddressesForAccount", reflect.TypeOf((*MockStore)(nil).GetActiveProxyClusterAddressesForAccount), ctx, accountID)
|
||||
}
|
||||
|
||||
// GetActiveProxyVersions mocks base method.
|
||||
func (m *MockStore) GetActiveProxyVersions(ctx context.Context, clusterAddr string) ([]string, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetActiveProxyVersions", ctx, clusterAddr)
|
||||
ret0, _ := ret[0].([]string)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetActiveProxyVersions indicates an expected call of GetActiveProxyVersions.
|
||||
func (mr *MockStoreMockRecorder) GetActiveProxyVersions(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveProxyVersions", reflect.TypeOf((*MockStore)(nil).GetActiveProxyVersions), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// GetAgentNetworkAccessLogSessions mocks base method.
|
||||
func (m *MockStore) GetAgentNetworkAccessLogSessions(ctx context.Context, lockStrength LockingStrength, accountID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error) {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
Reference in New Issue
Block a user