mirror of
https://github.com/netbirdio/netbird.git
synced 2026-09-29 18:19:07 +02:00
[management, proxy] Enforce strict base64url decoding for JWT validation (#7554)
This commit is contained in:
@@ -3,9 +3,11 @@ package auth
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
@@ -52,6 +54,22 @@ func (s *SessionStore) RegisterToken(ctx context.Context, token string, expiresA
|
||||
}
|
||||
|
||||
func hashToken(token string) string {
|
||||
sum := sha256.Sum256([]byte(token))
|
||||
sum := sha256.Sum256([]byte(canonicalizeToken(token)))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// canonicalizeToken re-encodes the JWT signature segment so noncanonical
|
||||
// spellings of the same signature map to one stable cache key.
|
||||
func canonicalizeToken(token string) string {
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != 3 {
|
||||
return token
|
||||
}
|
||||
|
||||
sig, err := base64.RawURLEncoding.DecodeString(parts[2])
|
||||
if err != nil {
|
||||
return token
|
||||
}
|
||||
|
||||
return parts[0] + "." + parts[1] + "." + base64.RawURLEncoding.EncodeToString(sig)
|
||||
}
|
||||
|
||||
@@ -2,10 +2,15 @@ package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -131,3 +136,49 @@ func TestHashToken_StableAndDoesNotLeak(t *testing.T) {
|
||||
assert.Len(t, a, 64, "sha256 hex must be 64 chars")
|
||||
assert.NotContains(t, a, "tokenA", "raw token must not appear in hash")
|
||||
}
|
||||
|
||||
func TestSessionStore_NoncanonicalSpellingIsRejectedAsReplay(t *testing.T) {
|
||||
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
require.NoError(t, err)
|
||||
|
||||
token := jwt.NewWithClaims(jwt.SigningMethodRS256, jwt.MapClaims{
|
||||
"sub": "user",
|
||||
"exp": time.Now().Add(time.Hour).Unix(),
|
||||
})
|
||||
canonical, err := token.SignedString(privateKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
parts := strings.Split(canonical, ".")
|
||||
require.Len(t, parts, 3)
|
||||
|
||||
// A 256-byte RSA signature (256 mod 3 == 1) leaves unused bits in the final
|
||||
// base64url character; flip one without changing the decoded signature.
|
||||
const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_"
|
||||
last := strings.IndexByte(alphabet, parts[2][len(parts[2])-1])
|
||||
require.GreaterOrEqual(t, last, 0)
|
||||
require.Equal(t, 0, last&3, "unexpected canonical RSA signature encoding")
|
||||
|
||||
equivalentSig := parts[2][:len(parts[2])-1] + string(alphabet[last|1])
|
||||
equivalent := parts[0] + "." + parts[1] + "." + equivalentSig
|
||||
require.NotEqual(t, canonical, equivalent, "spellings must differ as strings")
|
||||
|
||||
// Same decoded signature bytes, so they verify as the same JWT.
|
||||
canonicalSig, err := base64.RawURLEncoding.DecodeString(parts[2])
|
||||
require.NoError(t, err)
|
||||
altSig, err := base64.RawURLEncoding.DecodeString(equivalentSig)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, canonicalSig, altSig, "spellings must decode to identical signature bytes")
|
||||
|
||||
// The replay-cache key must be identical for both spellings.
|
||||
assert.Equal(t, hashToken(canonical), hashToken(equivalent),
|
||||
"noncanonical spelling must map to the same replay-cache key")
|
||||
|
||||
s := newTestSessionStore(t)
|
||||
ctx := context.Background()
|
||||
exp := time.Now().Add(time.Hour)
|
||||
|
||||
require.NoError(t, s.RegisterToken(ctx, canonical, exp), "first claim should succeed")
|
||||
err = s.RegisterToken(ctx, equivalent, exp)
|
||||
require.Error(t, err, "alternate spelling must be treated as a replay")
|
||||
assert.ErrorIs(t, err, ErrTokenAlreadyUsed)
|
||||
}
|
||||
|
||||
+1
-1
@@ -66,7 +66,7 @@ func ValidateSessionJWT(tokenString, domain string, publicKey ed25519.PublicKey)
|
||||
return nil, fmt.Errorf("unexpected signing method: %v", t.Header["alg"])
|
||||
}
|
||||
return publicKey, nil
|
||||
}, jwt.WithAudience(domain), jwt.WithIssuer(SessionJWTIssuer))
|
||||
}, jwt.WithAudience(domain), jwt.WithIssuer(SessionJWTIssuer), jwt.WithStrictDecoding())
|
||||
if err != nil {
|
||||
return "", "", "", nil, nil, fmt.Errorf("parse token: %w", err)
|
||||
}
|
||||
|
||||
@@ -18,7 +18,6 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
@@ -217,8 +216,8 @@ func (v *Validator) ValidateAndParse(ctx context.Context, token string) (*jwt.To
|
||||
jwt.WithAudience(v.audienceList...),
|
||||
jwt.WithIssuer(v.issuer),
|
||||
jwt.WithIssuedAt(),
|
||||
jwt.WithStrictDecoding(),
|
||||
)
|
||||
|
||||
// Check if there was an error in parsing...
|
||||
if err != nil {
|
||||
err = fmt.Errorf("%w: %s", errTokenParsing, err)
|
||||
|
||||
Reference in New Issue
Block a user