mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
+110
-34
@@ -2,6 +2,7 @@ package cryptokeys
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strconv"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
@@ -9,16 +10,14 @@ import (
|
||||
|
||||
"cdr.dev/slog"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
// never represents the maximum value for a time.Duration.
|
||||
const never = 1<<63 - 1
|
||||
|
||||
// DBCache implements Keycache for callers with access to the database.
|
||||
type DBCache struct {
|
||||
// dbCache implements Keycache for callers with access to the database.
|
||||
type dbCache struct {
|
||||
db database.Store
|
||||
feature database.CryptoKeyFeature
|
||||
logger slog.Logger
|
||||
@@ -34,18 +33,34 @@ type DBCache struct {
|
||||
closed bool
|
||||
}
|
||||
|
||||
type DBCacheOption func(*DBCache)
|
||||
type DBCacheOption func(*dbCache)
|
||||
|
||||
func WithDBCacheClock(clock quartz.Clock) DBCacheOption {
|
||||
return func(d *DBCache) {
|
||||
return func(d *dbCache) {
|
||||
d.clock = clock
|
||||
}
|
||||
}
|
||||
|
||||
// NewDBCache creates a new DBCache. Close should be called to
|
||||
// NewSigningCache creates a new DBCache. Close should be called to
|
||||
// release resources associated with its internal timer.
|
||||
func NewDBCache(logger slog.Logger, db database.Store, feature database.CryptoKeyFeature, opts ...func(*DBCache)) *DBCache {
|
||||
d := &DBCache{
|
||||
func NewSigningCache(logger slog.Logger, db database.Store, feature database.CryptoKeyFeature, opts ...func(*dbCache)) (SigningKeycache, error) {
|
||||
if !isSigningKeyFeature(feature) {
|
||||
return nil, ErrInvalidFeature
|
||||
}
|
||||
|
||||
return newDBCache(logger, db, feature, opts...), nil
|
||||
}
|
||||
|
||||
func NewEncryptionCache(logger slog.Logger, db database.Store, feature database.CryptoKeyFeature, opts ...func(*dbCache)) (EncryptionKeycache, error) {
|
||||
if !isEncryptionKeyFeature(feature) {
|
||||
return nil, ErrInvalidFeature
|
||||
}
|
||||
|
||||
return newDBCache(logger, db, feature, opts...), nil
|
||||
}
|
||||
|
||||
func newDBCache(logger slog.Logger, db database.Store, feature database.CryptoKeyFeature, opts ...func(*dbCache)) *dbCache {
|
||||
d := &dbCache{
|
||||
db: db,
|
||||
feature: feature,
|
||||
clock: quartz.NewReal(),
|
||||
@@ -56,23 +71,61 @@ func NewDBCache(logger slog.Logger, db database.Store, feature database.CryptoKe
|
||||
opt(d)
|
||||
}
|
||||
|
||||
// Initialize the timer. This will get properly initialized the first time we fetch.
|
||||
d.timer = d.clock.AfterFunc(never, d.clear)
|
||||
|
||||
return d
|
||||
}
|
||||
|
||||
// Verifying returns the CryptoKey with the given sequence number, provided that
|
||||
func (d *dbCache) EncryptingKey(ctx context.Context) (id string, key interface{}, err error) {
|
||||
if !isEncryptionKeyFeature(d.feature) {
|
||||
return "", nil, ErrInvalidFeature
|
||||
}
|
||||
|
||||
return d.latest(ctx)
|
||||
}
|
||||
|
||||
func (d *dbCache) DecryptingKey(ctx context.Context, id string) (key interface{}, err error) {
|
||||
if !isEncryptionKeyFeature(d.feature) {
|
||||
return nil, ErrInvalidFeature
|
||||
}
|
||||
|
||||
return d.sequence(ctx, id)
|
||||
}
|
||||
|
||||
func (d *dbCache) SigningKey(ctx context.Context) (id string, key interface{}, err error) {
|
||||
if !isSigningKeyFeature(d.feature) {
|
||||
return "", nil, ErrInvalidFeature
|
||||
}
|
||||
|
||||
return d.latest(ctx)
|
||||
}
|
||||
|
||||
func (d *dbCache) VerifyingKey(ctx context.Context, id string) (key interface{}, err error) {
|
||||
if !isSigningKeyFeature(d.feature) {
|
||||
return nil, ErrInvalidFeature
|
||||
}
|
||||
|
||||
return d.sequence(ctx, id)
|
||||
}
|
||||
|
||||
// sequence returns the CryptoKey with the given sequence number, provided that
|
||||
// it is neither deleted nor has breached its deletion date. It should only be
|
||||
// used for verifying or decrypting payloads. To sign/encrypt call Signing.
|
||||
func (d *DBCache) Verifying(ctx context.Context, sequence int32) (codersdk.CryptoKey, error) {
|
||||
func (d *dbCache) sequence(ctx context.Context, id string) (interface{}, error) {
|
||||
sequence, err := strconv.ParseInt(id, 10, 32)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("expecting sequence number got %q: %w", id, err)
|
||||
}
|
||||
|
||||
d.keysMu.RLock()
|
||||
if d.closed {
|
||||
d.keysMu.RUnlock()
|
||||
return codersdk.CryptoKey{}, ErrClosed
|
||||
return nil, ErrClosed
|
||||
}
|
||||
|
||||
now := d.clock.Now()
|
||||
key, ok := d.keys[sequence]
|
||||
key, ok := d.keys[int32(sequence)]
|
||||
d.keysMu.RUnlock()
|
||||
if ok {
|
||||
return checkKey(key, now)
|
||||
@@ -82,35 +135,35 @@ func (d *DBCache) Verifying(ctx context.Context, sequence int32) (codersdk.Crypt
|
||||
defer d.keysMu.Unlock()
|
||||
|
||||
if d.closed {
|
||||
return codersdk.CryptoKey{}, ErrClosed
|
||||
return nil, ErrClosed
|
||||
}
|
||||
|
||||
key, ok = d.keys[sequence]
|
||||
key, ok = d.keys[int32(sequence)]
|
||||
if ok {
|
||||
return checkKey(key, now)
|
||||
}
|
||||
|
||||
err := d.fetch(ctx)
|
||||
err = d.fetch(ctx)
|
||||
if err != nil {
|
||||
return codersdk.CryptoKey{}, xerrors.Errorf("fetch: %w", err)
|
||||
return nil, xerrors.Errorf("fetch: %w", err)
|
||||
}
|
||||
|
||||
key, ok = d.keys[sequence]
|
||||
key, ok = d.keys[int32(sequence)]
|
||||
if !ok {
|
||||
return codersdk.CryptoKey{}, ErrKeyNotFound
|
||||
return nil, ErrKeyNotFound
|
||||
}
|
||||
|
||||
return checkKey(key, now)
|
||||
}
|
||||
|
||||
// Signing returns the latest valid key for signing. A valid key is one that is
|
||||
// latest returns the latest valid key for signing. A valid key is one that is
|
||||
// both past its start time and before its deletion time.
|
||||
func (d *DBCache) Signing(ctx context.Context) (codersdk.CryptoKey, error) {
|
||||
func (d *dbCache) latest(ctx context.Context) (string, interface{}, error) {
|
||||
d.keysMu.RLock()
|
||||
|
||||
if d.closed {
|
||||
d.keysMu.RUnlock()
|
||||
return codersdk.CryptoKey{}, ErrClosed
|
||||
return "", nil, ErrClosed
|
||||
}
|
||||
|
||||
latest := d.latestKey
|
||||
@@ -118,31 +171,31 @@ func (d *DBCache) Signing(ctx context.Context) (codersdk.CryptoKey, error) {
|
||||
|
||||
now := d.clock.Now()
|
||||
if latest.CanSign(now) {
|
||||
return db2sdk.CryptoKey(latest), nil
|
||||
return idSecret(latest)
|
||||
}
|
||||
|
||||
d.keysMu.Lock()
|
||||
defer d.keysMu.Unlock()
|
||||
|
||||
if d.closed {
|
||||
return codersdk.CryptoKey{}, ErrClosed
|
||||
return "", nil, ErrClosed
|
||||
}
|
||||
|
||||
if d.latestKey.CanSign(now) {
|
||||
return db2sdk.CryptoKey(d.latestKey), nil
|
||||
return idSecret(d.latestKey)
|
||||
}
|
||||
|
||||
// Refetch all keys for this feature so we can find the latest valid key.
|
||||
err := d.fetch(ctx)
|
||||
if err != nil {
|
||||
return codersdk.CryptoKey{}, xerrors.Errorf("fetch: %w", err)
|
||||
return "", nil, xerrors.Errorf("fetch: %w", err)
|
||||
}
|
||||
|
||||
return db2sdk.CryptoKey(d.latestKey), nil
|
||||
return idSecret(d.latestKey)
|
||||
}
|
||||
|
||||
// clear invalidates the cache. This forces the subsequent call to fetch fresh keys.
|
||||
func (d *DBCache) clear() {
|
||||
func (d *dbCache) clear() {
|
||||
now := d.clock.Now("DBCache", "clear")
|
||||
d.keysMu.Lock()
|
||||
defer d.keysMu.Unlock()
|
||||
@@ -158,7 +211,7 @@ func (d *DBCache) clear() {
|
||||
|
||||
// fetch fetches all keys for the given feature and determines the latest key.
|
||||
// It must be called while holding the keysMu lock.
|
||||
func (d *DBCache) fetch(ctx context.Context) error {
|
||||
func (d *dbCache) fetch(ctx context.Context) error {
|
||||
keys, err := d.db.GetCryptoKeysByFeature(ctx, d.feature)
|
||||
if err != nil {
|
||||
return xerrors.Errorf("get crypto keys by feature: %w", err)
|
||||
@@ -189,22 +242,45 @@ func (d *DBCache) fetch(ctx context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func checkKey(key database.CryptoKey, now time.Time) (codersdk.CryptoKey, error) {
|
||||
func checkKey(key database.CryptoKey, now time.Time) (interface{}, error) {
|
||||
if !key.CanVerify(now) {
|
||||
return codersdk.CryptoKey{}, ErrKeyInvalid
|
||||
return nil, ErrKeyInvalid
|
||||
}
|
||||
|
||||
return db2sdk.CryptoKey(key), nil
|
||||
return key.DecodeString()
|
||||
}
|
||||
|
||||
func (d *DBCache) Close() {
|
||||
func (d *dbCache) Close() error {
|
||||
d.keysMu.Lock()
|
||||
defer d.keysMu.Unlock()
|
||||
|
||||
if d.closed {
|
||||
return
|
||||
return nil
|
||||
}
|
||||
|
||||
d.timer.Stop()
|
||||
d.closed = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func isEncryptionKeyFeature(feature database.CryptoKeyFeature) bool {
|
||||
return feature == database.CryptoKeyFeatureWorkspaceApps
|
||||
}
|
||||
|
||||
func isSigningKeyFeature(feature database.CryptoKeyFeature) bool {
|
||||
switch feature {
|
||||
case database.CryptoKeyFeatureTailnetResume, database.CryptoKeyFeatureOidcConvert:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func idSecret(k database.CryptoKey) (string, interface{}, error) {
|
||||
key, err := k.DecodeString()
|
||||
if err != nil {
|
||||
return "", nil, xerrors.Errorf("decode key: %w", err)
|
||||
}
|
||||
|
||||
return strconv.FormatInt(int64(k.Sequence), 10), key, nil
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package cryptokeys
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"strconv"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -11,13 +12,12 @@ import (
|
||||
"cdr.dev/slog/sloggers/slogtest"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/coderd/database/dbmock"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
func Test_Verifying(t *testing.T) {
|
||||
func Test_version(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("HitsCache", func(t *testing.T) {
|
||||
@@ -35,7 +35,7 @@ func Test_Verifying(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 32,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
}
|
||||
@@ -44,13 +44,13 @@ func Test_Verifying(t *testing.T) {
|
||||
32: expectedKey,
|
||||
}
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
k.keys = cache
|
||||
|
||||
got, err := k.Verifying(ctx, 32)
|
||||
secret, err := k.sequence(ctx, keyID(expectedKey))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(expectedKey), got)
|
||||
require.Equal(t, decodedSecret(t, expectedKey), secret)
|
||||
})
|
||||
|
||||
t.Run("MissesCache", func(t *testing.T) {
|
||||
@@ -69,20 +69,19 @@ func Test_Verifying(t *testing.T) {
|
||||
Sequence: 33,
|
||||
StartsAt: clock.Now(),
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
}
|
||||
|
||||
mockDB.EXPECT().GetCryptoKeysByFeature(ctx, database.CryptoKeyFeatureWorkspaceApps).Return([]database.CryptoKey{expectedKey}, nil)
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
|
||||
got, err := k.Verifying(ctx, 33)
|
||||
got, err := k.sequence(ctx, keyID(expectedKey))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(expectedKey), got)
|
||||
require.Equal(t, db2sdk.CryptoKey(expectedKey), db2sdk.CryptoKey(k.latestKey))
|
||||
require.Equal(t, decodedSecret(t, expectedKey), got)
|
||||
})
|
||||
|
||||
t.Run("InvalidCachedKey", func(t *testing.T) {
|
||||
@@ -101,7 +100,7 @@ func Test_Verifying(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 32,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
DeletesAt: sql.NullTime{
|
||||
@@ -111,11 +110,11 @@ func Test_Verifying(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
k.keys = cache
|
||||
|
||||
_, err := k.Verifying(ctx, 32)
|
||||
_, err := k.sequence(ctx, "32")
|
||||
require.ErrorIs(t, err, ErrKeyInvalid)
|
||||
})
|
||||
|
||||
@@ -134,7 +133,7 @@ func Test_Verifying(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 32,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
DeletesAt: sql.NullTime{
|
||||
@@ -144,15 +143,15 @@ func Test_Verifying(t *testing.T) {
|
||||
}
|
||||
mockDB.EXPECT().GetCryptoKeysByFeature(ctx, database.CryptoKeyFeatureWorkspaceApps).Return([]database.CryptoKey{invalidKey}, nil)
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
|
||||
_, err := k.Verifying(ctx, 32)
|
||||
_, err := k.sequence(ctx, keyID(invalidKey))
|
||||
require.ErrorIs(t, err, ErrKeyInvalid)
|
||||
})
|
||||
}
|
||||
|
||||
func Test_Signing(t *testing.T) {
|
||||
func Test_latest(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("HitsCache", func(t *testing.T) {
|
||||
@@ -170,19 +169,20 @@ func Test_Signing(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 32,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now(),
|
||||
}
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
|
||||
k.latestKey = latestKey
|
||||
|
||||
got, err := k.Signing(ctx)
|
||||
id, secret, err := k.latest(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(latestKey), got)
|
||||
require.Equal(t, keyID(latestKey), id)
|
||||
require.Equal(t, decodedSecret(t, latestKey), secret)
|
||||
})
|
||||
|
||||
t.Run("InvalidCachedKey", func(t *testing.T) {
|
||||
@@ -200,7 +200,7 @@ func Test_Signing(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 33,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now(),
|
||||
@@ -210,7 +210,7 @@ func Test_Signing(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 32,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now().Add(-time.Hour),
|
||||
@@ -222,13 +222,14 @@ func Test_Signing(t *testing.T) {
|
||||
|
||||
mockDB.EXPECT().GetCryptoKeysByFeature(ctx, database.CryptoKeyFeatureWorkspaceApps).Return([]database.CryptoKey{latestKey}, nil)
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
k.latestKey = invalidKey
|
||||
|
||||
got, err := k.Signing(ctx)
|
||||
id, secret, err := k.latest(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(latestKey), got)
|
||||
require.Equal(t, keyID(latestKey), id)
|
||||
require.Equal(t, decodedSecret(t, latestKey), secret)
|
||||
})
|
||||
|
||||
t.Run("UsesActiveKey", func(t *testing.T) {
|
||||
@@ -246,7 +247,7 @@ func Test_Signing(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 32,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now().Add(time.Hour),
|
||||
@@ -256,7 +257,7 @@ func Test_Signing(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 33,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now(),
|
||||
@@ -264,12 +265,13 @@ func Test_Signing(t *testing.T) {
|
||||
|
||||
mockDB.EXPECT().GetCryptoKeysByFeature(ctx, database.CryptoKeyFeatureWorkspaceApps).Return([]database.CryptoKey{inactiveKey, activeKey}, nil)
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
|
||||
got, err := k.Signing(ctx)
|
||||
id, secret, err := k.latest(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(activeKey), got)
|
||||
require.Equal(t, keyID(activeKey), id)
|
||||
require.Equal(t, decodedSecret(t, activeKey), secret)
|
||||
})
|
||||
|
||||
t.Run("NoValidKeys", func(t *testing.T) {
|
||||
@@ -287,7 +289,7 @@ func Test_Signing(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 32,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now().Add(time.Hour),
|
||||
@@ -297,7 +299,7 @@ func Test_Signing(t *testing.T) {
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 33,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now().Add(-time.Hour),
|
||||
@@ -309,10 +311,10 @@ func Test_Signing(t *testing.T) {
|
||||
|
||||
mockDB.EXPECT().GetCryptoKeysByFeature(ctx, database.CryptoKeyFeatureWorkspaceApps).Return([]database.CryptoKey{inactiveKey, invalidKey}, nil)
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
|
||||
_, err := k.Signing(ctx)
|
||||
_, _, err := k.latest(ctx)
|
||||
require.ErrorIs(t, err, ErrKeyInvalid)
|
||||
})
|
||||
}
|
||||
@@ -331,14 +333,14 @@ func Test_clear(t *testing.T) {
|
||||
logger = slogtest.Make(t, nil)
|
||||
)
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
|
||||
activeKey := database.CryptoKey{
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 33,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now(),
|
||||
@@ -346,7 +348,7 @@ func Test_clear(t *testing.T) {
|
||||
|
||||
mockDB.EXPECT().GetCryptoKeysByFeature(ctx, database.CryptoKeyFeatureWorkspaceApps).Return([]database.CryptoKey{activeKey}, nil)
|
||||
|
||||
_, err := k.Signing(ctx)
|
||||
_, _, err := k.latest(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
dur, wait := clock.AdvanceNext()
|
||||
@@ -367,14 +369,14 @@ func Test_clear(t *testing.T) {
|
||||
logger = slogtest.Make(t, nil)
|
||||
)
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
|
||||
key := database.CryptoKey{
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 32,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now(),
|
||||
@@ -386,9 +388,10 @@ func Test_clear(t *testing.T) {
|
||||
// timer is reset and doesn't fire after another five minute.
|
||||
clock.Advance(time.Minute * 5)
|
||||
|
||||
latest, err := k.Signing(ctx)
|
||||
id, secret, err := k.latest(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(key), latest)
|
||||
require.Equal(t, keyID(key), id)
|
||||
require.Equal(t, decodedSecret(t, key), secret)
|
||||
|
||||
// Advancing the clock now should require 10 minutes
|
||||
// before the timer fires again.
|
||||
@@ -415,14 +418,14 @@ func Test_clear(t *testing.T) {
|
||||
|
||||
trap := clock.Trap().Now("clear")
|
||||
|
||||
k := NewDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
k := newDBCache(logger, mockDB, database.CryptoKeyFeatureWorkspaceApps, WithDBCacheClock(clock))
|
||||
defer k.Close()
|
||||
|
||||
key := database.CryptoKey{
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Sequence: 32,
|
||||
Secret: sql.NullString{
|
||||
String: "secret",
|
||||
String: mustGenerateKey(t),
|
||||
Valid: true,
|
||||
},
|
||||
StartsAt: clock.Now(),
|
||||
@@ -431,9 +434,10 @@ func Test_clear(t *testing.T) {
|
||||
mockDB.EXPECT().GetCryptoKeysByFeature(ctx, database.CryptoKeyFeatureWorkspaceApps).Return([]database.CryptoKey{key}, nil).Times(2)
|
||||
|
||||
// Move us past the initial timer.
|
||||
latest, err := k.Signing(ctx)
|
||||
id, secret, err := k.latest(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(key), latest)
|
||||
require.Equal(t, keyID(key), id)
|
||||
require.Equal(t, decodedSecret(t, key), secret)
|
||||
// Null these out so that we refetch.
|
||||
k.keys = nil
|
||||
k.latestKey = database.CryptoKey{}
|
||||
@@ -445,9 +449,10 @@ func Test_clear(t *testing.T) {
|
||||
call := trap.MustWait(ctx)
|
||||
|
||||
// Refetch keys.
|
||||
latest, err = k.Signing(ctx)
|
||||
id, secret, err = k.latest(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(key), latest)
|
||||
require.Equal(t, keyID(key), id)
|
||||
require.Equal(t, decodedSecret(t, key), secret)
|
||||
|
||||
// Let the rest of the timer function run.
|
||||
// It should see that we have refetched keys and
|
||||
@@ -465,3 +470,21 @@ func Test_clear(t *testing.T) {
|
||||
require.Equal(t, database.CryptoKey{}, k.latestKey)
|
||||
})
|
||||
}
|
||||
|
||||
func mustGenerateKey(t *testing.T) string {
|
||||
t.Helper()
|
||||
key, err := generateKey(64)
|
||||
require.NoError(t, err)
|
||||
return key
|
||||
}
|
||||
|
||||
func keyID(key database.CryptoKey) string {
|
||||
return strconv.FormatInt(int64(key.Sequence), 10)
|
||||
}
|
||||
|
||||
func decodedSecret(t *testing.T, key database.CryptoKey) []byte {
|
||||
t.Helper()
|
||||
decoded, err := key.DecodeString()
|
||||
require.NoError(t, err)
|
||||
return decoded
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package cryptokeys_test
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -10,7 +11,6 @@ import (
|
||||
|
||||
"github.com/coder/coder/v2/coderd/cryptokeys"
|
||||
"github.com/coder/coder/v2/coderd/database"
|
||||
"github.com/coder/coder/v2/coderd/database/db2sdk"
|
||||
"github.com/coder/coder/v2/coderd/database/dbgen"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
@@ -24,7 +24,7 @@ func TestMain(m *testing.M) {
|
||||
func TestDBKeyCache(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("Verifying", func(t *testing.T) {
|
||||
t.Run("VerifyingKey", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("HitsCache", func(t *testing.T) {
|
||||
@@ -38,17 +38,18 @@ func TestDBKeyCache(t *testing.T) {
|
||||
)
|
||||
|
||||
key := dbgen.CryptoKey(t, db, database.CryptoKey{
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Feature: database.CryptoKeyFeatureOidcConvert,
|
||||
Sequence: 1,
|
||||
StartsAt: clock.Now().UTC(),
|
||||
})
|
||||
|
||||
k := cryptokeys.NewDBCache(logger, db, database.CryptoKeyFeatureWorkspaceApps, cryptokeys.WithDBCacheClock(clock))
|
||||
k, err := cryptokeys.NewSigningCache(logger, db, database.CryptoKeyFeatureOidcConvert, cryptokeys.WithDBCacheClock(clock))
|
||||
require.NoError(t, err)
|
||||
defer k.Close()
|
||||
|
||||
got, err := k.Verifying(ctx, key.Sequence)
|
||||
got, err := k.VerifyingKey(ctx, keyID(key))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(key), got)
|
||||
require.Equal(t, decodedSecret(t, key), got)
|
||||
})
|
||||
|
||||
t.Run("NotFound", func(t *testing.T) {
|
||||
@@ -61,10 +62,11 @@ func TestDBKeyCache(t *testing.T) {
|
||||
logger = slogtest.Make(t, nil)
|
||||
)
|
||||
|
||||
k := cryptokeys.NewDBCache(logger, db, database.CryptoKeyFeatureWorkspaceApps, cryptokeys.WithDBCacheClock(clock))
|
||||
k, err := cryptokeys.NewSigningCache(logger, db, database.CryptoKeyFeatureOidcConvert, cryptokeys.WithDBCacheClock(clock))
|
||||
require.NoError(t, err)
|
||||
defer k.Close()
|
||||
|
||||
_, err := k.Verifying(ctx, 123)
|
||||
_, err = k.VerifyingKey(ctx, "123")
|
||||
require.ErrorIs(t, err, cryptokeys.ErrKeyNotFound)
|
||||
})
|
||||
})
|
||||
@@ -80,29 +82,31 @@ func TestDBKeyCache(t *testing.T) {
|
||||
)
|
||||
|
||||
_ = dbgen.CryptoKey(t, db, database.CryptoKey{
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Feature: database.CryptoKeyFeatureOidcConvert,
|
||||
Sequence: 10,
|
||||
StartsAt: clock.Now().UTC(),
|
||||
})
|
||||
|
||||
expectedKey := dbgen.CryptoKey(t, db, database.CryptoKey{
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Feature: database.CryptoKeyFeatureOidcConvert,
|
||||
Sequence: 12,
|
||||
StartsAt: clock.Now().UTC(),
|
||||
})
|
||||
|
||||
_ = dbgen.CryptoKey(t, db, database.CryptoKey{
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Feature: database.CryptoKeyFeatureOidcConvert,
|
||||
Sequence: 2,
|
||||
StartsAt: clock.Now().UTC(),
|
||||
})
|
||||
|
||||
k := cryptokeys.NewDBCache(logger, db, database.CryptoKeyFeatureWorkspaceApps, cryptokeys.WithDBCacheClock(clock))
|
||||
k, err := cryptokeys.NewSigningCache(logger, db, database.CryptoKeyFeatureOidcConvert, cryptokeys.WithDBCacheClock(clock))
|
||||
require.NoError(t, err)
|
||||
defer k.Close()
|
||||
|
||||
got, err := k.Signing(ctx)
|
||||
id, key, err := k.SigningKey(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(expectedKey), got)
|
||||
require.Equal(t, keyID(expectedKey), id)
|
||||
require.Equal(t, decodedSecret(t, expectedKey), key)
|
||||
})
|
||||
|
||||
t.Run("Closed", func(t *testing.T) {
|
||||
@@ -116,28 +120,97 @@ func TestDBKeyCache(t *testing.T) {
|
||||
)
|
||||
|
||||
expectedKey := dbgen.CryptoKey(t, db, database.CryptoKey{
|
||||
Feature: database.CryptoKeyFeatureWorkspaceApps,
|
||||
Feature: database.CryptoKeyFeatureOidcConvert,
|
||||
Sequence: 10,
|
||||
StartsAt: clock.Now(),
|
||||
})
|
||||
|
||||
k := cryptokeys.NewDBCache(logger, db, database.CryptoKeyFeatureWorkspaceApps, cryptokeys.WithDBCacheClock(clock))
|
||||
k, err := cryptokeys.NewSigningCache(logger, db, database.CryptoKeyFeatureOidcConvert, cryptokeys.WithDBCacheClock(clock))
|
||||
require.NoError(t, err)
|
||||
defer k.Close()
|
||||
|
||||
got, err := k.Signing(ctx)
|
||||
id, key, err := k.SigningKey(ctx)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(expectedKey), got)
|
||||
require.Equal(t, keyID(expectedKey), id)
|
||||
require.Equal(t, decodedSecret(t, expectedKey), key)
|
||||
|
||||
got, err = k.Verifying(ctx, expectedKey.Sequence)
|
||||
key, err = k.VerifyingKey(ctx, keyID(expectedKey))
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, db2sdk.CryptoKey(expectedKey), got)
|
||||
require.Equal(t, decodedSecret(t, expectedKey), key)
|
||||
|
||||
k.Close()
|
||||
|
||||
_, err = k.Signing(ctx)
|
||||
_, _, err = k.SigningKey(ctx)
|
||||
require.ErrorIs(t, err, cryptokeys.ErrClosed)
|
||||
|
||||
_, err = k.Verifying(ctx, expectedKey.Sequence)
|
||||
_, err = k.VerifyingKey(ctx, keyID(expectedKey))
|
||||
require.ErrorIs(t, err, cryptokeys.ErrClosed)
|
||||
})
|
||||
|
||||
t.Run("InvalidSigningFeature", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var (
|
||||
db, _ = dbtestutil.NewDB(t)
|
||||
clock = quartz.NewMock(t)
|
||||
logger = slogtest.Make(t, nil)
|
||||
ctx = testutil.Context(t, testutil.WaitShort)
|
||||
)
|
||||
|
||||
_, err := cryptokeys.NewSigningCache(logger, db, database.CryptoKeyFeatureWorkspaceApps, cryptokeys.WithDBCacheClock(clock))
|
||||
require.ErrorIs(t, err, cryptokeys.ErrInvalidFeature)
|
||||
|
||||
// Instantiate a signing cache and try to use it as an encryption cache.
|
||||
sc, err := cryptokeys.NewSigningCache(logger, db, database.CryptoKeyFeatureOidcConvert, cryptokeys.WithDBCacheClock(clock))
|
||||
require.NoError(t, err)
|
||||
defer sc.Close()
|
||||
|
||||
ec, ok := sc.(cryptokeys.EncryptionKeycache)
|
||||
require.True(t, ok)
|
||||
_, _, err = ec.EncryptingKey(ctx)
|
||||
require.ErrorIs(t, err, cryptokeys.ErrInvalidFeature)
|
||||
|
||||
_, err = ec.DecryptingKey(ctx, "123")
|
||||
require.ErrorIs(t, err, cryptokeys.ErrInvalidFeature)
|
||||
})
|
||||
|
||||
t.Run("InvalidEncryptionFeature", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
var (
|
||||
db, _ = dbtestutil.NewDB(t)
|
||||
clock = quartz.NewMock(t)
|
||||
logger = slogtest.Make(t, nil)
|
||||
ctx = testutil.Context(t, testutil.WaitShort)
|
||||
)
|
||||
|
||||
_, err := cryptokeys.NewEncryptionCache(logger, db, database.CryptoKeyFeatureOidcConvert, cryptokeys.WithDBCacheClock(clock))
|
||||
require.ErrorIs(t, err, cryptokeys.ErrInvalidFeature)
|
||||
|
||||
// Instantiate an encryption cache and try to use it as a signing cache.
|
||||
ec, err := cryptokeys.NewEncryptionCache(logger, db, database.CryptoKeyFeatureWorkspaceApps, cryptokeys.WithDBCacheClock(clock))
|
||||
require.NoError(t, err)
|
||||
defer ec.Close()
|
||||
|
||||
sc, ok := ec.(cryptokeys.SigningKeycache)
|
||||
require.True(t, ok)
|
||||
_, _, err = sc.SigningKey(ctx)
|
||||
require.ErrorIs(t, err, cryptokeys.ErrInvalidFeature)
|
||||
|
||||
_, err = sc.VerifyingKey(ctx, "123")
|
||||
require.ErrorIs(t, err, cryptokeys.ErrInvalidFeature)
|
||||
})
|
||||
}
|
||||
|
||||
func keyID(key database.CryptoKey) string {
|
||||
return strconv.FormatInt(int64(key.Sequence), 10)
|
||||
}
|
||||
|
||||
func decodedSecret(t *testing.T, key database.CryptoKey) []byte {
|
||||
t.Helper()
|
||||
|
||||
secret, err := key.DecodeString()
|
||||
require.NoError(t, err)
|
||||
|
||||
return secret
|
||||
}
|
||||
|
||||
@@ -0,0 +1,2 @@
|
||||
// Package cryptokeys provides an abstraction for fetching internally used cryptographic keys mainly for JWT signing and verification.
|
||||
package cryptokeys
|
||||
@@ -2,20 +2,40 @@ package cryptokeys
|
||||
|
||||
import (
|
||||
"context"
|
||||
"io"
|
||||
|
||||
"golang.org/x/xerrors"
|
||||
|
||||
"github.com/coder/coder/v2/codersdk"
|
||||
)
|
||||
|
||||
var (
|
||||
ErrKeyNotFound = xerrors.New("key not found")
|
||||
ErrKeyInvalid = xerrors.New("key is invalid for use")
|
||||
ErrClosed = xerrors.New("closed")
|
||||
ErrKeyNotFound = xerrors.New("key not found")
|
||||
ErrKeyInvalid = xerrors.New("key is invalid for use")
|
||||
ErrClosed = xerrors.New("closed")
|
||||
ErrInvalidFeature = xerrors.New("invalid feature for this operation")
|
||||
)
|
||||
|
||||
// Keycache provides an abstraction for fetching signing keys.
|
||||
type Keycache interface {
|
||||
Signing(ctx context.Context) (codersdk.CryptoKey, error)
|
||||
Verifying(ctx context.Context, sequence int32) (codersdk.CryptoKey, error)
|
||||
type EncryptionKeycache interface {
|
||||
// EncryptingKey returns the latest valid key for encrypting payloads. A valid
|
||||
// key is one that is both past its start time and before its deletion time.
|
||||
EncryptingKey(ctx context.Context) (id string, key interface{}, err error)
|
||||
// DecryptingKey returns the key with the provided id which maps to its sequence
|
||||
// number. The key is valid for decryption as long as it is not deleted or past
|
||||
// its deletion date. We must allow for keys prior to their start time to
|
||||
// account for clock skew between peers (one key may be past its start time on
|
||||
// one machine while another is not).
|
||||
DecryptingKey(ctx context.Context, id string) (key interface{}, err error)
|
||||
io.Closer
|
||||
}
|
||||
|
||||
type SigningKeycache interface {
|
||||
// SigningKey returns the latest valid key for signing. A valid key is one
|
||||
// that is both past its start time and before its deletion time.
|
||||
SigningKey(ctx context.Context) (id string, key interface{}, err error)
|
||||
// VerifyingKey returns the key with the provided id which should map to its
|
||||
// sequence number. The key is valid for verifying as long as it is not deleted
|
||||
// or past its deletion date. We must allow for keys prior to their start time
|
||||
// to account for clock skew between peers (one key may be past its start time
|
||||
// on one machine while another is not).
|
||||
VerifyingKey(ctx context.Context, id string) (key interface{}, err error)
|
||||
io.Closer
|
||||
}
|
||||
|
||||
@@ -227,9 +227,9 @@ func (k *rotator) rotateKey(ctx context.Context, tx database.Store, key database
|
||||
func generateNewSecret(feature database.CryptoKeyFeature) (string, error) {
|
||||
switch feature {
|
||||
case database.CryptoKeyFeatureWorkspaceApps:
|
||||
return generateKey(96)
|
||||
case database.CryptoKeyFeatureOidcConvert:
|
||||
return generateKey(32)
|
||||
case database.CryptoKeyFeatureOidcConvert:
|
||||
return generateKey(64)
|
||||
case database.CryptoKeyFeatureTailnetResume:
|
||||
return generateKey(64)
|
||||
}
|
||||
|
||||
@@ -588,9 +588,9 @@ func requireKey(t *testing.T, key database.CryptoKey, feature database.CryptoKey
|
||||
|
||||
switch key.Feature {
|
||||
case database.CryptoKeyFeatureOidcConvert:
|
||||
require.Len(t, secret, 32)
|
||||
require.Len(t, secret, 64)
|
||||
case database.CryptoKeyFeatureWorkspaceApps:
|
||||
require.Len(t, secret, 96)
|
||||
require.Len(t, secret, 32)
|
||||
case database.CryptoKeyFeatureTailnetResume:
|
||||
require.Len(t, secret, 64)
|
||||
default:
|
||||
|
||||
Reference in New Issue
Block a user