feat: add jwt pkg (#14928)

- Adds a `jwtutils` package to be shared amongst the various
packages in the codebase that make use of JWTs. It's intended to help us
standardize on one library instead of some implementations using
`go-jose` and others using `golang-jwt`.

The main reason we're converging on `go-jose` is due to its support for
JWEs, `golang-jwt` also has a repo to handle it but it doesn't look
maintained: https://github.com/golang-jwt/jwe
This commit is contained in:
Jon Ayers
2024-10-03 21:09:52 -05:00
committed by GitHub
parent 50d9206950
commit 68ec532ca7
13 changed files with 1001 additions and 122 deletions
+121
View File
@@ -0,0 +1,121 @@
package jwtutils
import (
"context"
"encoding/json"
"time"
"github.com/go-jose/go-jose/v4"
"github.com/go-jose/go-jose/v4/jwt"
"golang.org/x/xerrors"
)
const (
encryptKeyAlgo = jose.A256GCMKW
encryptContentAlgo = jose.A256GCM
)
type EncryptKeyProvider interface {
EncryptingKey(ctx context.Context) (id string, key interface{}, err error)
}
type DecryptKeyProvider interface {
DecryptingKey(ctx context.Context, id string) (key interface{}, err error)
}
// Encrypt encrypts a token and returns it as a string.
func Encrypt(ctx context.Context, e EncryptKeyProvider, claims Claims) (string, error) {
id, key, err := e.EncryptingKey(ctx)
if err != nil {
return "", xerrors.Errorf("get signing key: %w", err)
}
encrypter, err := jose.NewEncrypter(
encryptContentAlgo,
jose.Recipient{
Algorithm: encryptKeyAlgo,
Key: key,
},
&jose.EncrypterOptions{
Compression: jose.DEFLATE,
ExtraHeaders: map[jose.HeaderKey]interface{}{
keyIDHeaderKey: id,
},
},
)
if err != nil {
return "", xerrors.Errorf("initialize encrypter: %w", err)
}
payload, err := json.Marshal(claims)
if err != nil {
return "", xerrors.Errorf("marshal payload: %w", err)
}
encrypted, err := encrypter.Encrypt(payload)
if err != nil {
return "", xerrors.Errorf("encrypt: %w", err)
}
compact, err := encrypted.CompactSerialize()
if err != nil {
return "", xerrors.Errorf("compact serialize: %w", err)
}
return compact, nil
}
// DecryptOptions are options for decrypting a JWE.
type DecryptOptions struct {
RegisteredClaims jwt.Expected
KeyAlgorithm jose.KeyAlgorithm
ContentEncryptionAlgorithm jose.ContentEncryption
}
// Decrypt decrypts the token using the provided key. It unmarshals into the provided claims.
func Decrypt(ctx context.Context, d DecryptKeyProvider, token string, claims Claims, opts ...func(*DecryptOptions)) error {
options := DecryptOptions{
RegisteredClaims: jwt.Expected{
Time: time.Now(),
},
KeyAlgorithm: encryptKeyAlgo,
ContentEncryptionAlgorithm: encryptContentAlgo,
}
for _, opt := range opts {
opt(&options)
}
object, err := jose.ParseEncrypted(token,
[]jose.KeyAlgorithm{options.KeyAlgorithm},
[]jose.ContentEncryption{options.ContentEncryptionAlgorithm},
)
if err != nil {
return xerrors.Errorf("parse jwe: %w", err)
}
if object.Header.Algorithm != string(encryptKeyAlgo) {
return xerrors.Errorf("expected JWE algorithm to be %q, got %q", encryptKeyAlgo, object.Header.Algorithm)
}
kid := object.Header.KeyID
if kid == "" {
return xerrors.Errorf("expected %q header to be a string", keyIDHeaderKey)
}
key, err := d.DecryptingKey(ctx, kid)
if err != nil {
return xerrors.Errorf("key with id %q: %w", kid, err)
}
decrypted, err := object.Decrypt(key)
if err != nil {
return xerrors.Errorf("decrypt: %w", err)
}
if err := json.Unmarshal(decrypted, &claims); err != nil {
return xerrors.Errorf("unmarshal: %w", err)
}
return claims.Validate(options.RegisteredClaims)
}
+127
View File
@@ -0,0 +1,127 @@
package jwtutils
import (
"context"
"encoding/json"
"time"
"github.com/go-jose/go-jose/v4"
"github.com/go-jose/go-jose/v4/jwt"
"golang.org/x/xerrors"
)
const (
keyIDHeaderKey = "kid"
)
// Claims defines the payload for a JWT. Most callers
// should embed jwt.Claims
type Claims interface {
Validate(jwt.Expected) error
}
const (
signingAlgo = jose.HS512
)
type SigningKeyProvider interface {
SigningKey(ctx context.Context) (id string, key interface{}, err error)
}
type VerifyKeyProvider interface {
VerifyingKey(ctx context.Context, id string) (key interface{}, err error)
}
// Sign signs a token and returns it as a string.
func Sign(ctx context.Context, s SigningKeyProvider, claims Claims) (string, error) {
id, key, err := s.SigningKey(ctx)
if err != nil {
return "", xerrors.Errorf("get signing key: %w", err)
}
signer, err := jose.NewSigner(jose.SigningKey{
Algorithm: signingAlgo,
Key: key,
}, &jose.SignerOptions{
ExtraHeaders: map[jose.HeaderKey]interface{}{
keyIDHeaderKey: id,
},
})
if err != nil {
return "", xerrors.Errorf("new signer: %w", err)
}
payload, err := json.Marshal(claims)
if err != nil {
return "", xerrors.Errorf("marshal claims: %w", err)
}
signed, err := signer.Sign(payload)
if err != nil {
return "", xerrors.Errorf("sign payload: %w", err)
}
compact, err := signed.CompactSerialize()
if err != nil {
return "", xerrors.Errorf("compact serialize: %w", err)
}
return compact, nil
}
// VerifyOptions are options for verifying a JWT.
type VerifyOptions struct {
RegisteredClaims jwt.Expected
SignatureAlgorithm jose.SignatureAlgorithm
}
// Verify verifies that a token was signed by the provided key. It unmarshals into the provided claims.
func Verify(ctx context.Context, v VerifyKeyProvider, token string, claims Claims, opts ...func(*VerifyOptions)) error {
options := VerifyOptions{
RegisteredClaims: jwt.Expected{
Time: time.Now(),
},
SignatureAlgorithm: signingAlgo,
}
for _, opt := range opts {
opt(&options)
}
object, err := jose.ParseSigned(token, []jose.SignatureAlgorithm{options.SignatureAlgorithm})
if err != nil {
return xerrors.Errorf("parse JWS: %w", err)
}
if len(object.Signatures) != 1 {
return xerrors.New("expected 1 signature")
}
signature := object.Signatures[0]
if signature.Header.Algorithm != string(signingAlgo) {
return xerrors.Errorf("expected JWS algorithm to be %q, got %q", signingAlgo, object.Signatures[0].Header.Algorithm)
}
kid := signature.Header.KeyID
if kid == "" {
return xerrors.Errorf("expected %q header to be a string", keyIDHeaderKey)
}
key, err := v.VerifyingKey(ctx, kid)
if err != nil {
return xerrors.Errorf("key with id %q: %w", kid, err)
}
payload, err := object.Verify(key)
if err != nil {
return xerrors.Errorf("verify payload: %w", err)
}
err = json.Unmarshal(payload, &claims)
if err != nil {
return xerrors.Errorf("unmarshal payload: %w", err)
}
return claims.Validate(options.RegisteredClaims)
}
+436
View File
@@ -0,0 +1,436 @@
package jwtutils_test
import (
"context"
"crypto/rand"
"testing"
"time"
"github.com/go-jose/go-jose/v4"
"github.com/go-jose/go-jose/v4/jwt"
"github.com/google/uuid"
"github.com/stretchr/testify/require"
"cdr.dev/slog/sloggers/slogtest"
"github.com/coder/coder/v2/coderd/cryptokeys"
"github.com/coder/coder/v2/coderd/database"
"github.com/coder/coder/v2/coderd/database/dbgen"
"github.com/coder/coder/v2/coderd/database/dbtestutil"
"github.com/coder/coder/v2/coderd/jwtutils"
"github.com/coder/coder/v2/testutil"
)
func TestClaims(t *testing.T) {
t.Parallel()
type tokenType struct {
Name string
KeySize int
Sign bool
}
types := []tokenType{
{
Name: "JWE",
Sign: false,
KeySize: 32,
},
{
Name: "JWS",
Sign: true,
KeySize: 64,
},
}
type testcase struct {
name string
claims jwtutils.Claims
expectedClaims jwt.Expected
expectedErr error
}
cases := []testcase{
{
name: "OK",
claims: jwt.Claims{
Issuer: "coder",
Subject: "user@coder.com",
Audience: jwt.Audience{"coder"},
Expiry: jwt.NewNumericDate(time.Now().Add(time.Hour)),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
},
},
{
name: "WrongIssuer",
claims: jwt.Claims{
Issuer: "coder",
Subject: "user@coder.com",
Audience: jwt.Audience{"coder"},
Expiry: jwt.NewNumericDate(time.Now().Add(time.Hour)),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
},
expectedClaims: jwt.Expected{
Issuer: "coder2",
},
expectedErr: jwt.ErrInvalidIssuer,
},
{
name: "WrongSubject",
claims: jwt.Claims{
Issuer: "coder",
Subject: "user@coder.com",
Audience: jwt.Audience{"coder"},
Expiry: jwt.NewNumericDate(time.Now().Add(time.Hour)),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
},
expectedClaims: jwt.Expected{
Subject: "user2@coder.com",
},
expectedErr: jwt.ErrInvalidSubject,
},
{
name: "WrongAudience",
claims: jwt.Claims{
Issuer: "coder",
Subject: "user@coder.com",
Audience: jwt.Audience{"coder"},
Expiry: jwt.NewNumericDate(time.Now().Add(time.Hour)),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
},
},
{
name: "Expired",
claims: jwt.Claims{
Issuer: "coder",
Subject: "user@coder.com",
Audience: jwt.Audience{"coder"},
Expiry: jwt.NewNumericDate(time.Now().Add(time.Minute)),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now()),
},
expectedClaims: jwt.Expected{
Time: time.Now().Add(time.Minute * 3),
},
expectedErr: jwt.ErrExpired,
},
{
name: "IssuedInFuture",
claims: jwt.Claims{
Issuer: "coder",
Subject: "user@coder.com",
Audience: jwt.Audience{"coder"},
Expiry: jwt.NewNumericDate(time.Now().Add(time.Minute)),
IssuedAt: jwt.NewNumericDate(time.Now()),
},
expectedClaims: jwt.Expected{
Time: time.Now().Add(-time.Minute * 3),
},
expectedErr: jwt.ErrIssuedInTheFuture,
},
{
name: "IsBefore",
claims: jwt.Claims{
Issuer: "coder",
Subject: "user@coder.com",
Audience: jwt.Audience{"coder"},
Expiry: jwt.NewNumericDate(time.Now().Add(time.Hour)),
IssuedAt: jwt.NewNumericDate(time.Now()),
NotBefore: jwt.NewNumericDate(time.Now().Add(time.Minute * 5)),
},
expectedClaims: jwt.Expected{
Time: time.Now().Add(time.Minute * 3),
},
expectedErr: jwt.ErrNotValidYet,
},
}
for _, tt := range types {
tt := tt
t.Run(tt.Name, func(t *testing.T) {
t.Parallel()
for _, c := range cases {
c := c
t.Run(c.name, func(t *testing.T) {
t.Parallel()
var (
ctx = testutil.Context(t, testutil.WaitShort)
key = newKey(t, tt.KeySize)
token string
err error
)
if tt.Sign {
token, err = jwtutils.Sign(ctx, key, c.claims)
} else {
token, err = jwtutils.Encrypt(ctx, key, c.claims)
}
require.NoError(t, err)
var actual jwt.Claims
if tt.Sign {
err = jwtutils.Verify(ctx, key, token, &actual, withVerifyExpected(c.expectedClaims))
} else {
err = jwtutils.Decrypt(ctx, key, token, &actual, withDecryptExpected(c.expectedClaims))
}
if c.expectedErr != nil {
require.ErrorIs(t, err, c.expectedErr)
} else {
require.NoError(t, err)
require.Equal(t, c.claims, actual)
}
})
}
})
}
}
func TestJWS(t *testing.T) {
t.Parallel()
t.Run("WrongSignatureAlgorithm", func(t *testing.T) {
t.Parallel()
ctx := testutil.Context(t, testutil.WaitShort)
key := newKey(t, 64)
token, err := jwtutils.Sign(ctx, key, jwt.Claims{})
require.NoError(t, err)
var actual testClaims
err = jwtutils.Verify(ctx, key, token, &actual, withSignatureAlgorithm(jose.HS256))
require.Error(t, err)
})
t.Run("CustomClaims", func(t *testing.T) {
t.Parallel()
var (
ctx = testutil.Context(t, testutil.WaitShort)
key = newKey(t, 64)
)
expected := testClaims{
MyClaim: "my_value",
}
token, err := jwtutils.Sign(ctx, key, expected)
require.NoError(t, err)
var actual testClaims
err = jwtutils.Verify(ctx, key, token, &actual, withVerifyExpected(jwt.Expected{}))
require.NoError(t, err)
require.Equal(t, expected, actual)
})
t.Run("WithKeycache", func(t *testing.T) {
t.Parallel()
var (
ctx = testutil.Context(t, testutil.WaitShort)
db, _ = dbtestutil.NewDB(t)
_ = dbgen.CryptoKey(t, db, database.CryptoKey{
Feature: database.CryptoKeyFeatureOidcConvert,
StartsAt: time.Now(),
})
log = slogtest.Make(t, nil)
)
cache, err := cryptokeys.NewSigningCache(log, db, database.CryptoKeyFeatureOidcConvert)
require.NoError(t, err)
claims := testClaims{
MyClaim: "my_value",
Claims: jwt.Claims{
Expiry: jwt.NewNumericDate(time.Now().Add(time.Hour)),
},
}
token, err := jwtutils.Sign(ctx, cache, claims)
require.NoError(t, err)
var actual testClaims
err = jwtutils.Verify(ctx, cache, token, &actual)
require.NoError(t, err)
require.Equal(t, claims, actual)
})
}
func TestJWE(t *testing.T) {
t.Parallel()
t.Run("WrongKeyAlgorithm", func(t *testing.T) {
t.Parallel()
var (
ctx = testutil.Context(t, testutil.WaitShort)
key = newKey(t, 32)
)
token, err := jwtutils.Encrypt(ctx, key, jwt.Claims{})
require.NoError(t, err)
var actual testClaims
err = jwtutils.Decrypt(ctx, key, token, &actual, withKeyAlgorithm(jose.A128GCMKW))
require.Error(t, err)
})
t.Run("WrongContentyEncryption", func(t *testing.T) {
t.Parallel()
var (
ctx = testutil.Context(t, testutil.WaitShort)
key = newKey(t, 32)
)
token, err := jwtutils.Encrypt(ctx, key, jwt.Claims{})
require.NoError(t, err)
var actual testClaims
err = jwtutils.Decrypt(ctx, key, token, &actual, withContentEncryptionAlgorithm(jose.A128GCM))
require.Error(t, err)
})
t.Run("CustomClaims", func(t *testing.T) {
t.Parallel()
var (
ctx = testutil.Context(t, testutil.WaitShort)
key = newKey(t, 32)
)
expected := testClaims{
MyClaim: "my_value",
}
token, err := jwtutils.Encrypt(ctx, key, expected)
require.NoError(t, err)
var actual testClaims
err = jwtutils.Decrypt(ctx, key, token, &actual, withDecryptExpected(jwt.Expected{}))
require.NoError(t, err)
require.Equal(t, expected, actual)
})
t.Run("WithKeycache", func(t *testing.T) {
t.Parallel()
var (
ctx = testutil.Context(t, testutil.WaitShort)
db, _ = dbtestutil.NewDB(t)
_ = dbgen.CryptoKey(t, db, database.CryptoKey{
Feature: database.CryptoKeyFeatureWorkspaceApps,
StartsAt: time.Now(),
})
log = slogtest.Make(t, nil)
)
cache, err := cryptokeys.NewEncryptionCache(log, db, database.CryptoKeyFeatureWorkspaceApps)
require.NoError(t, err)
claims := testClaims{
MyClaim: "my_value",
Claims: jwt.Claims{
Expiry: jwt.NewNumericDate(time.Now().Add(time.Hour)),
},
}
token, err := jwtutils.Encrypt(ctx, cache, claims)
require.NoError(t, err)
var actual testClaims
err = jwtutils.Decrypt(ctx, cache, token, &actual)
require.NoError(t, err)
require.Equal(t, claims, actual)
})
}
func generateSecret(t *testing.T, keySize int) []byte {
t.Helper()
b := make([]byte, keySize)
_, err := rand.Read(b)
require.NoError(t, err)
return b
}
type testClaims struct {
MyClaim string `json:"my_claim"`
jwt.Claims
}
func withDecryptExpected(e jwt.Expected) func(*jwtutils.DecryptOptions) {
return func(opts *jwtutils.DecryptOptions) {
opts.RegisteredClaims = e
}
}
func withVerifyExpected(e jwt.Expected) func(*jwtutils.VerifyOptions) {
return func(opts *jwtutils.VerifyOptions) {
opts.RegisteredClaims = e
}
}
func withSignatureAlgorithm(alg jose.SignatureAlgorithm) func(*jwtutils.VerifyOptions) {
return func(opts *jwtutils.VerifyOptions) {
opts.SignatureAlgorithm = alg
}
}
func withKeyAlgorithm(alg jose.KeyAlgorithm) func(*jwtutils.DecryptOptions) {
return func(opts *jwtutils.DecryptOptions) {
opts.KeyAlgorithm = alg
}
}
func withContentEncryptionAlgorithm(alg jose.ContentEncryption) func(*jwtutils.DecryptOptions) {
return func(opts *jwtutils.DecryptOptions) {
opts.ContentEncryptionAlgorithm = alg
}
}
type key struct {
t testing.TB
id string
secret []byte
}
func newKey(t *testing.T, size int) *key {
t.Helper()
id := uuid.New().String()
secret := generateSecret(t, size)
return &key{
t: t,
id: id,
secret: secret,
}
}
func (k *key) SigningKey(_ context.Context) (id string, key interface{}, err error) {
return k.id, k.secret, nil
}
func (k *key) VerifyingKey(_ context.Context, id string) (key interface{}, err error) {
k.t.Helper()
require.Equal(k.t, k.id, id)
return k.secret, nil
}
func (k *key) EncryptingKey(_ context.Context) (id string, key interface{}, err error) {
return k.id, k.secret, nil
}
func (k *key) DecryptingKey(_ context.Context, id string) (key interface{}, err error) {
k.t.Helper()
require.Equal(k.t, k.id, id)
return k.secret, nil
}