[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:
Bethuel Mmbaga
2026-09-29 18:29:55 +03:00
committed by GitHub
parent 7ff709f565
commit 30dd076b36
22 changed files with 676 additions and 281 deletions
+52 -19
View File
@@ -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
}
+26 -2
View File
@@ -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())
}