mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
feat: enable key rotation (#15066)
This PR contains the remaining logic necessary to hook up key rotation to the product.
This commit is contained in:
+26
-117
@@ -3,32 +3,23 @@ package tailnet
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"time"
|
||||
|
||||
"github.com/go-jose/go-jose/v3"
|
||||
"github.com/go-jose/go-jose/v4/jwt"
|
||||
"github.com/google/uuid"
|
||||
"golang.org/x/xerrors"
|
||||
"google.golang.org/protobuf/types/known/durationpb"
|
||||
"google.golang.org/protobuf/types/known/timestamppb"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/jwtutils"
|
||||
"github.com/coder/coder/v2/tailnet/proto"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
const (
|
||||
DefaultResumeTokenExpiry = 24 * time.Hour
|
||||
|
||||
resumeTokenSigningAlgorithm = jose.HS512
|
||||
)
|
||||
|
||||
// resumeTokenSigningKeyID is a fixed key ID for the resume token signing key.
|
||||
// If/when we add support for multiple keys (e.g. key rotation), this will move
|
||||
// to the database instead.
|
||||
var resumeTokenSigningKeyID = uuid.MustParse("97166747-9309-4d7f-9071-a230e257c2a4")
|
||||
|
||||
// NewInsecureTestResumeTokenProvider returns a ResumeTokenProvider that uses a
|
||||
// random key with short expiry for testing purposes. If any errors occur while
|
||||
// generating the key, the function panics.
|
||||
@@ -37,12 +28,15 @@ func NewInsecureTestResumeTokenProvider() ResumeTokenProvider {
|
||||
if err != nil {
|
||||
panic(err)
|
||||
}
|
||||
return NewResumeTokenKeyProvider(key, quartz.NewReal(), time.Hour)
|
||||
return NewResumeTokenKeyProvider(jwtutils.StaticKey{
|
||||
ID: uuid.New().String(),
|
||||
Key: key[:],
|
||||
}, quartz.NewReal(), time.Hour)
|
||||
}
|
||||
|
||||
type ResumeTokenProvider interface {
|
||||
GenerateResumeToken(peerID uuid.UUID) (*proto.RefreshResumeTokenResponse, error)
|
||||
VerifyResumeToken(token string) (uuid.UUID, error)
|
||||
GenerateResumeToken(ctx context.Context, peerID uuid.UUID) (*proto.RefreshResumeTokenResponse, error)
|
||||
VerifyResumeToken(ctx context.Context, token string) (uuid.UUID, error)
|
||||
}
|
||||
|
||||
type ResumeTokenSigningKey [64]byte
|
||||
@@ -56,104 +50,37 @@ func GenerateResumeTokenSigningKey() (ResumeTokenSigningKey, error) {
|
||||
return key, nil
|
||||
}
|
||||
|
||||
type ResumeTokenSigningKeyDatabaseStore interface {
|
||||
GetCoordinatorResumeTokenSigningKey(ctx context.Context) (string, error)
|
||||
UpsertCoordinatorResumeTokenSigningKey(ctx context.Context, key string) error
|
||||
}
|
||||
|
||||
// ResumeTokenSigningKeyFromDatabase retrieves the coordinator resume token
|
||||
// signing key from the database. If the key is not found, a new key is
|
||||
// generated and inserted into the database.
|
||||
func ResumeTokenSigningKeyFromDatabase(ctx context.Context, db ResumeTokenSigningKeyDatabaseStore) (ResumeTokenSigningKey, error) {
|
||||
var resumeTokenKey ResumeTokenSigningKey
|
||||
resumeTokenKeyStr, err := db.GetCoordinatorResumeTokenSigningKey(ctx)
|
||||
if err != nil && !xerrors.Is(err, sql.ErrNoRows) {
|
||||
return resumeTokenKey, xerrors.Errorf("get coordinator resume token key: %w", err)
|
||||
}
|
||||
if decoded, err := hex.DecodeString(resumeTokenKeyStr); err != nil || len(decoded) != len(resumeTokenKey) {
|
||||
newKey, err := GenerateResumeTokenSigningKey()
|
||||
if err != nil {
|
||||
return resumeTokenKey, xerrors.Errorf("generate fresh coordinator resume token key: %w", err)
|
||||
}
|
||||
|
||||
resumeTokenKeyStr = hex.EncodeToString(newKey[:])
|
||||
err = db.UpsertCoordinatorResumeTokenSigningKey(ctx, resumeTokenKeyStr)
|
||||
if err != nil {
|
||||
return resumeTokenKey, xerrors.Errorf("insert freshly generated coordinator resume token key to database: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
resumeTokenKeyBytes, err := hex.DecodeString(resumeTokenKeyStr)
|
||||
if err != nil {
|
||||
return resumeTokenKey, xerrors.Errorf("decode coordinator resume token key from database: %w", err)
|
||||
}
|
||||
if len(resumeTokenKeyBytes) != len(resumeTokenKey) {
|
||||
return resumeTokenKey, xerrors.Errorf("coordinator resume token key in database is not the correct length, expect %d got %d", len(resumeTokenKey), len(resumeTokenKeyBytes))
|
||||
}
|
||||
copy(resumeTokenKey[:], resumeTokenKeyBytes)
|
||||
if resumeTokenKey == [64]byte{} {
|
||||
return resumeTokenKey, xerrors.Errorf("coordinator resume token key in database is empty")
|
||||
}
|
||||
return resumeTokenKey, nil
|
||||
}
|
||||
|
||||
type ResumeTokenKeyProvider struct {
|
||||
key ResumeTokenSigningKey
|
||||
key jwtutils.SigningKeyManager
|
||||
clock quartz.Clock
|
||||
expiry time.Duration
|
||||
}
|
||||
|
||||
func NewResumeTokenKeyProvider(key ResumeTokenSigningKey, clock quartz.Clock, expiry time.Duration) ResumeTokenProvider {
|
||||
func NewResumeTokenKeyProvider(key jwtutils.SigningKeyManager, clock quartz.Clock, expiry time.Duration) ResumeTokenProvider {
|
||||
if expiry <= 0 {
|
||||
expiry = DefaultResumeTokenExpiry
|
||||
}
|
||||
return ResumeTokenKeyProvider{
|
||||
key: key,
|
||||
clock: clock,
|
||||
expiry: DefaultResumeTokenExpiry,
|
||||
expiry: expiry,
|
||||
}
|
||||
}
|
||||
|
||||
type resumeTokenPayload struct {
|
||||
PeerID uuid.UUID `json:"sub"`
|
||||
Expiry int64 `json:"exp"`
|
||||
}
|
||||
|
||||
func (p ResumeTokenKeyProvider) GenerateResumeToken(peerID uuid.UUID) (*proto.RefreshResumeTokenResponse, error) {
|
||||
func (p ResumeTokenKeyProvider) GenerateResumeToken(ctx context.Context, peerID uuid.UUID) (*proto.RefreshResumeTokenResponse, error) {
|
||||
exp := p.clock.Now().Add(p.expiry)
|
||||
payload := resumeTokenPayload{
|
||||
PeerID: peerID,
|
||||
Expiry: exp.Unix(),
|
||||
}
|
||||
payloadBytes, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("marshal payload to JSON: %w", err)
|
||||
payload := jwtutils.RegisteredClaims{
|
||||
Subject: peerID.String(),
|
||||
Expiry: jwt.NewNumericDate(exp),
|
||||
}
|
||||
|
||||
signer, err := jose.NewSigner(jose.SigningKey{
|
||||
Algorithm: resumeTokenSigningAlgorithm,
|
||||
Key: p.key[:],
|
||||
}, &jose.SignerOptions{
|
||||
ExtraHeaders: map[jose.HeaderKey]interface{}{
|
||||
"kid": resumeTokenSigningKeyID.String(),
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("create signer: %w", err)
|
||||
}
|
||||
|
||||
signedObject, err := signer.Sign(payloadBytes)
|
||||
token, err := jwtutils.Sign(ctx, p.key, payload)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("sign payload: %w", err)
|
||||
}
|
||||
|
||||
serialized, err := signedObject.CompactSerialize()
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("serialize JWS: %w", err)
|
||||
}
|
||||
|
||||
return &proto.RefreshResumeTokenResponse{
|
||||
Token: serialized,
|
||||
Token: token,
|
||||
RefreshIn: durationpb.New(p.expiry / 2),
|
||||
ExpiresAt: timestamppb.New(exp),
|
||||
}, nil
|
||||
@@ -162,35 +89,17 @@ func (p ResumeTokenKeyProvider) GenerateResumeToken(peerID uuid.UUID) (*proto.Re
|
||||
// VerifyResumeToken parses a signed tailnet resume token with the given key and
|
||||
// returns the payload. If the token is invalid or expired, an error is
|
||||
// returned.
|
||||
func (p ResumeTokenKeyProvider) VerifyResumeToken(str string) (uuid.UUID, error) {
|
||||
object, err := jose.ParseSigned(str)
|
||||
func (p ResumeTokenKeyProvider) VerifyResumeToken(ctx context.Context, str string) (uuid.UUID, error) {
|
||||
var tok jwt.Claims
|
||||
err := jwtutils.Verify(ctx, p.key, str, &tok, jwtutils.WithVerifyExpected(jwt.Expected{
|
||||
Time: p.clock.Now(),
|
||||
}))
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf("parse JWS: %w", err)
|
||||
return uuid.Nil, xerrors.Errorf("verify payload: %w", err)
|
||||
}
|
||||
if len(object.Signatures) != 1 {
|
||||
return uuid.Nil, xerrors.New("expected 1 signature")
|
||||
}
|
||||
if object.Signatures[0].Header.Algorithm != string(resumeTokenSigningAlgorithm) {
|
||||
return uuid.Nil, xerrors.Errorf("expected token signing algorithm to be %q, got %q", resumeTokenSigningAlgorithm, object.Signatures[0].Header.Algorithm)
|
||||
}
|
||||
if object.Signatures[0].Header.KeyID != resumeTokenSigningKeyID.String() {
|
||||
return uuid.Nil, xerrors.Errorf("expected token key ID to be %q, got %q", resumeTokenSigningKeyID, object.Signatures[0].Header.KeyID)
|
||||
}
|
||||
|
||||
output, err := object.Verify(p.key[:])
|
||||
parsed, err := uuid.Parse(tok.Subject)
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf("verify JWS: %w", err)
|
||||
return uuid.Nil, xerrors.Errorf("parse peerID from token: %w", err)
|
||||
}
|
||||
|
||||
var tok resumeTokenPayload
|
||||
err = json.Unmarshal(output, &tok)
|
||||
if err != nil {
|
||||
return uuid.Nil, xerrors.Errorf("unmarshal payload: %w", err)
|
||||
}
|
||||
exp := time.Unix(tok.Expiry, 0)
|
||||
if exp.Before(p.clock.Now()) {
|
||||
return uuid.Nil, xerrors.New("signed resume token expired")
|
||||
}
|
||||
|
||||
return tok.PeerID, nil
|
||||
return parsed, nil
|
||||
}
|
||||
|
||||
+34
-116
@@ -1,117 +1,20 @@
|
||||
package tailnet_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"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/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/coder/coder/v2/coderd/database/dbmock"
|
||||
"github.com/coder/coder/v2/coderd/database/dbtestutil"
|
||||
"github.com/coder/coder/v2/coderd/jwtutils"
|
||||
"github.com/coder/coder/v2/tailnet"
|
||||
"github.com/coder/coder/v2/testutil"
|
||||
"github.com/coder/quartz"
|
||||
)
|
||||
|
||||
func TestResumeTokenSigningKeyFromDatabase(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
assertRandomKey := func(t *testing.T, key tailnet.ResumeTokenSigningKey) {
|
||||
t.Helper()
|
||||
assert.NotEqual(t, tailnet.ResumeTokenSigningKey{}, key, "key should not be empty")
|
||||
assert.NotEqualValues(t, [64]byte{1}, key, "key should not be all 1s")
|
||||
}
|
||||
|
||||
t.Run("GenerateRetrieve", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db, _ := dbtestutil.NewDB(t)
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
key1, err := tailnet.ResumeTokenSigningKeyFromDatabase(ctx, db)
|
||||
require.NoError(t, err)
|
||||
assertRandomKey(t, key1)
|
||||
|
||||
key2, err := tailnet.ResumeTokenSigningKeyFromDatabase(ctx, db)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, key1, key2, "keys should not be different")
|
||||
})
|
||||
|
||||
t.Run("GetError", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db := dbmock.NewMockStore(gomock.NewController(t))
|
||||
db.EXPECT().GetCoordinatorResumeTokenSigningKey(gomock.Any()).Return("", assert.AnError)
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
_, err := tailnet.ResumeTokenSigningKeyFromDatabase(ctx, db)
|
||||
require.ErrorIs(t, err, assert.AnError)
|
||||
})
|
||||
|
||||
t.Run("UpsertError", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db := dbmock.NewMockStore(gomock.NewController(t))
|
||||
db.EXPECT().GetCoordinatorResumeTokenSigningKey(gomock.Any()).Return("", nil)
|
||||
db.EXPECT().UpsertCoordinatorResumeTokenSigningKey(gomock.Any(), gomock.Any()).Return(assert.AnError)
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
_, err := tailnet.ResumeTokenSigningKeyFromDatabase(ctx, db)
|
||||
require.ErrorIs(t, err, assert.AnError)
|
||||
})
|
||||
|
||||
t.Run("DecodeErrorShouldRegenerate", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db := dbmock.NewMockStore(gomock.NewController(t))
|
||||
db.EXPECT().GetCoordinatorResumeTokenSigningKey(gomock.Any()).Return("invalid", nil)
|
||||
|
||||
var storedKey tailnet.ResumeTokenSigningKey
|
||||
db.EXPECT().UpsertCoordinatorResumeTokenSigningKey(gomock.Any(), gomock.Any()).Do(func(_ context.Context, value string) error {
|
||||
keyBytes, err := hex.DecodeString(value)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, keyBytes, len(storedKey))
|
||||
copy(storedKey[:], keyBytes)
|
||||
return nil
|
||||
})
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
key, err := tailnet.ResumeTokenSigningKeyFromDatabase(ctx, db)
|
||||
require.NoError(t, err)
|
||||
assertRandomKey(t, key)
|
||||
require.Equal(t, storedKey, key, "key should match stored value")
|
||||
})
|
||||
|
||||
t.Run("LengthErrorShouldRegenerate", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db := dbmock.NewMockStore(gomock.NewController(t))
|
||||
db.EXPECT().GetCoordinatorResumeTokenSigningKey(gomock.Any()).Return("deadbeef", nil)
|
||||
db.EXPECT().UpsertCoordinatorResumeTokenSigningKey(gomock.Any(), gomock.Any()).Return(nil)
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
key, err := tailnet.ResumeTokenSigningKeyFromDatabase(ctx, db)
|
||||
require.NoError(t, err)
|
||||
assertRandomKey(t, key)
|
||||
})
|
||||
|
||||
t.Run("EmptyError", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
db := dbmock.NewMockStore(gomock.NewController(t))
|
||||
emptyKey := hex.EncodeToString(make([]byte, 64))
|
||||
db.EXPECT().GetCoordinatorResumeTokenSigningKey(gomock.Any()).Return(emptyKey, nil)
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
_, err := tailnet.ResumeTokenSigningKeyFromDatabase(ctx, db)
|
||||
require.ErrorContains(t, err, "is empty")
|
||||
})
|
||||
}
|
||||
|
||||
func TestResumeTokenKeyProvider(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
@@ -121,17 +24,18 @@ func TestResumeTokenKeyProvider(t *testing.T) {
|
||||
t.Run("OK", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
id := uuid.New()
|
||||
clock := quartz.NewMock(t)
|
||||
provider := tailnet.NewResumeTokenKeyProvider(key, clock, tailnet.DefaultResumeTokenExpiry)
|
||||
token, err := provider.GenerateResumeToken(id)
|
||||
provider := tailnet.NewResumeTokenKeyProvider(newKeySigner(key), clock, tailnet.DefaultResumeTokenExpiry)
|
||||
token, err := provider.GenerateResumeToken(ctx, id)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, token)
|
||||
require.NotEmpty(t, token.Token)
|
||||
require.Equal(t, tailnet.DefaultResumeTokenExpiry/2, token.RefreshIn.AsDuration())
|
||||
require.WithinDuration(t, clock.Now().Add(tailnet.DefaultResumeTokenExpiry), token.ExpiresAt.AsTime(), time.Second)
|
||||
|
||||
gotID, err := provider.VerifyResumeToken(token.Token)
|
||||
gotID, err := provider.VerifyResumeToken(ctx, token.Token)
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, id, gotID)
|
||||
})
|
||||
@@ -139,43 +43,57 @@ func TestResumeTokenKeyProvider(t *testing.T) {
|
||||
t.Run("Expired", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
id := uuid.New()
|
||||
clock := quartz.NewMock(t)
|
||||
provider := tailnet.NewResumeTokenKeyProvider(key, clock, tailnet.DefaultResumeTokenExpiry)
|
||||
token, err := provider.GenerateResumeToken(id)
|
||||
provider := tailnet.NewResumeTokenKeyProvider(newKeySigner(key), clock, tailnet.DefaultResumeTokenExpiry)
|
||||
token, err := provider.GenerateResumeToken(ctx, id)
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, token)
|
||||
require.NotEmpty(t, token.Token)
|
||||
require.Equal(t, tailnet.DefaultResumeTokenExpiry/2, token.RefreshIn.AsDuration())
|
||||
require.WithinDuration(t, clock.Now().Add(tailnet.DefaultResumeTokenExpiry), token.ExpiresAt.AsTime(), time.Second)
|
||||
|
||||
// Advance time past expiry
|
||||
_ = clock.Advance(tailnet.DefaultResumeTokenExpiry + time.Second)
|
||||
// Advance time past expiry. Account for leeway.
|
||||
_ = clock.Advance(tailnet.DefaultResumeTokenExpiry + time.Second*61)
|
||||
|
||||
_, err = provider.VerifyResumeToken(token.Token)
|
||||
require.ErrorContains(t, err, "expired")
|
||||
_, err = provider.VerifyResumeToken(ctx, token.Token)
|
||||
require.Error(t, err)
|
||||
require.ErrorIs(t, err, jwt.ErrExpired)
|
||||
})
|
||||
|
||||
t.Run("InvalidToken", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
provider := tailnet.NewResumeTokenKeyProvider(key, quartz.NewMock(t), tailnet.DefaultResumeTokenExpiry)
|
||||
_, err := provider.VerifyResumeToken("invalid")
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
provider := tailnet.NewResumeTokenKeyProvider(newKeySigner(key), quartz.NewMock(t), tailnet.DefaultResumeTokenExpiry)
|
||||
_, err := provider.VerifyResumeToken(ctx, "invalid")
|
||||
require.ErrorContains(t, err, "parse JWS")
|
||||
})
|
||||
|
||||
t.Run("VerifyError", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
ctx := testutil.Context(t, testutil.WaitShort)
|
||||
// Generate a resume token with a different key
|
||||
otherKey, err := tailnet.GenerateResumeTokenSigningKey()
|
||||
require.NoError(t, err)
|
||||
otherProvider := tailnet.NewResumeTokenKeyProvider(otherKey, quartz.NewMock(t), tailnet.DefaultResumeTokenExpiry)
|
||||
token, err := otherProvider.GenerateResumeToken(uuid.New())
|
||||
otherSigner := newKeySigner(otherKey)
|
||||
otherProvider := tailnet.NewResumeTokenKeyProvider(otherSigner, quartz.NewMock(t), tailnet.DefaultResumeTokenExpiry)
|
||||
token, err := otherProvider.GenerateResumeToken(ctx, uuid.New())
|
||||
require.NoError(t, err)
|
||||
|
||||
provider := tailnet.NewResumeTokenKeyProvider(key, quartz.NewMock(t), tailnet.DefaultResumeTokenExpiry)
|
||||
_, err = provider.VerifyResumeToken(token.Token)
|
||||
require.ErrorContains(t, err, "verify JWS")
|
||||
signer := newKeySigner(key)
|
||||
signer.ID = otherSigner.ID
|
||||
provider := tailnet.NewResumeTokenKeyProvider(signer, quartz.NewMock(t), tailnet.DefaultResumeTokenExpiry)
|
||||
_, err = provider.VerifyResumeToken(ctx, token.Token)
|
||||
require.ErrorIs(t, err, jose.ErrCryptoFailure)
|
||||
})
|
||||
}
|
||||
|
||||
func newKeySigner(key tailnet.ResumeTokenSigningKey) jwtutils.StaticKey {
|
||||
return jwtutils.StaticKey{
|
||||
ID: "123",
|
||||
Key: key[:],
|
||||
}
|
||||
}
|
||||
|
||||
+1
-1
@@ -177,7 +177,7 @@ func (s *DRPCService) RefreshResumeToken(ctx context.Context, _ *proto.RefreshRe
|
||||
return nil, xerrors.New("no Stream ID")
|
||||
}
|
||||
|
||||
res, err := s.ResumeTokenProvider.GenerateResumeToken(streamID.ID)
|
||||
res, err := s.ResumeTokenProvider.GenerateResumeToken(ctx, streamID.ID)
|
||||
if err != nil {
|
||||
return nil, xerrors.Errorf("generate resume token: %w", err)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user