Skip to content

Commit 021db7a

Browse files
committed
fix: harden auth service and reduce KMS load
1 parent 14a4d8a commit 021db7a

9 files changed

Lines changed: 331 additions & 28 deletions

File tree

internal/dolthubauth/kms.go

Lines changed: 111 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -5,9 +5,12 @@ import (
55
"crypto/aes"
66
"crypto/cipher"
77
"crypto/rand"
8+
"crypto/sha256"
89
"encoding/json"
910
"fmt"
1011
"strings"
12+
"sync"
13+
"time"
1114

1215
kms "cloud.google.com/go/kms/apiv1"
1316
"cloud.google.com/go/kms/apiv1/kmspb"
@@ -42,8 +45,9 @@ func (c realKMSEnvelopeClient) GetCryptoKey(ctx context.Context, req *kmspb.GetC
4245
// GCPKMSEnvelopeCipher encrypts credentials with a local DEK and wraps the DEK
4346
// with Google Cloud KMS.
4447
type GCPKMSEnvelopeCipher struct {
45-
client kmsEnvelopeClient
46-
keyName string
48+
client kmsEnvelopeClient
49+
keyName string
50+
dekCache *kmsDEKCache
4751
}
4852

4953
type kmsEnvelopePayload struct {
@@ -80,8 +84,9 @@ func newGCPKMSEnvelopeCipherWithClient(client kmsEnvelopeClient, keyName string)
8084
return nil, fmt.Errorf("kms client is required")
8185
}
8286
return &GCPKMSEnvelopeCipher{
83-
client: client,
84-
keyName: keyName,
87+
client: client,
88+
keyName: keyName,
89+
dekCache: newKMSDEKCache(5*time.Minute, 1024),
8590
}, nil
8691
}
8792

@@ -160,14 +165,19 @@ func (c *GCPKMSEnvelopeCipher) Decrypt(ctx context.Context, ciphertext []byte, k
160165
keyVersion = c.keyName
161166
}
162167

163-
unwrapped, err := c.client.Decrypt(ctx, &kmspb.DecryptRequest{
164-
Name: kmsEnvelopeCryptoKeyName(keyVersion, c.keyName),
165-
Ciphertext: payload.WrappedDEK,
166-
})
167-
if err != nil {
168-
return nil, fmt.Errorf("unwrap dek: %w", err)
168+
cacheKey := kmsEnvelopeDEKCacheKey(keyVersion, payload.WrappedDEK)
169+
dek, ok := c.dekCache.Get(cacheKey)
170+
if !ok {
171+
unwrapped, err := c.client.Decrypt(ctx, &kmspb.DecryptRequest{
172+
Name: kmsEnvelopeCryptoKeyName(keyVersion, c.keyName),
173+
Ciphertext: payload.WrappedDEK,
174+
})
175+
if err != nil {
176+
return nil, fmt.Errorf("unwrap dek: %w", err)
177+
}
178+
dek = append([]byte(nil), unwrapped.GetPlaintext()...)
179+
c.dekCache.Set(cacheKey, dek)
169180
}
170-
dek := append([]byte(nil), unwrapped.GetPlaintext()...)
171181
defer clearBytes(dek)
172182

173183
block, err := aes.NewCipher(dek)
@@ -202,6 +212,96 @@ func kmsEnvelopeAAD(keyVersion string) []byte {
202212
return []byte("wasteland:dolthub-auth:gcp-kms-envelope:" + strings.TrimSpace(keyVersion))
203213
}
204214

215+
func kmsEnvelopeDEKCacheKey(keyVersion string, wrappedDEK []byte) string {
216+
sum := sha256.Sum256(wrappedDEK)
217+
return strings.TrimSpace(keyVersion) + ":" + fmt.Sprintf("%x", sum[:])
218+
}
219+
220+
type kmsDEKCache struct {
221+
mu sync.Mutex
222+
ttl time.Duration
223+
maxEntries int
224+
entries map[string]kmsDEKCacheEntry
225+
}
226+
227+
type kmsDEKCacheEntry struct {
228+
dek []byte
229+
expiresAt time.Time
230+
usedAt time.Time
231+
}
232+
233+
func newKMSDEKCache(ttl time.Duration, maxEntries int) *kmsDEKCache {
234+
if ttl <= 0 || maxEntries <= 0 {
235+
return nil
236+
}
237+
return &kmsDEKCache{
238+
ttl: ttl,
239+
maxEntries: maxEntries,
240+
entries: make(map[string]kmsDEKCacheEntry),
241+
}
242+
}
243+
244+
func (c *kmsDEKCache) Get(key string) ([]byte, bool) {
245+
if c == nil {
246+
return nil, false
247+
}
248+
now := time.Now()
249+
c.mu.Lock()
250+
defer c.mu.Unlock()
251+
entry, ok := c.entries[key]
252+
if !ok {
253+
return nil, false
254+
}
255+
if now.After(entry.expiresAt) {
256+
clearBytes(entry.dek)
257+
delete(c.entries, key)
258+
return nil, false
259+
}
260+
entry.usedAt = now
261+
c.entries[key] = entry
262+
return append([]byte(nil), entry.dek...), true
263+
}
264+
265+
func (c *kmsDEKCache) Set(key string, dek []byte) {
266+
if c == nil || key == "" || len(dek) == 0 {
267+
return
268+
}
269+
now := time.Now()
270+
c.mu.Lock()
271+
defer c.mu.Unlock()
272+
if old, ok := c.entries[key]; ok {
273+
clearBytes(old.dek)
274+
}
275+
c.entries[key] = kmsDEKCacheEntry{
276+
dek: append([]byte(nil), dek...),
277+
expiresAt: now.Add(c.ttl),
278+
usedAt: now,
279+
}
280+
c.evictLocked(now)
281+
}
282+
283+
func (c *kmsDEKCache) evictLocked(now time.Time) {
284+
for key, entry := range c.entries {
285+
if now.After(entry.expiresAt) {
286+
clearBytes(entry.dek)
287+
delete(c.entries, key)
288+
}
289+
}
290+
for len(c.entries) > c.maxEntries {
291+
var oldestKey string
292+
var oldestAt time.Time
293+
for key, entry := range c.entries {
294+
if oldestKey == "" || entry.usedAt.Before(oldestAt) {
295+
oldestKey = key
296+
oldestAt = entry.usedAt
297+
}
298+
}
299+
entry := c.entries[oldestKey]
300+
clearBytes(entry.dek)
301+
delete(c.entries, oldestKey)
302+
}
303+
}
304+
205305
func clearBytes(buf []byte) {
206306
for i := range buf {
207307
buf[i] = 0

internal/dolthubauth/kms_test.go

Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ type fakeKMSEnvelopeClient struct {
1414
lastDecryptName string
1515
lastGetName string
1616
versionName string
17+
decryptCalls int
1718
checkErr error
1819
encryptErr error
1920
decryptErr error
@@ -34,6 +35,7 @@ func (f *fakeKMSEnvelopeClient) Decrypt(_ context.Context, req *kmspb.DecryptReq
3435
if f.decryptErr != nil {
3536
return nil, f.decryptErr
3637
}
38+
f.decryptCalls++
3739
f.lastDecryptName = req.GetName()
3840
if len(req.GetCiphertext()) < len("wrapped:") || string(req.GetCiphertext()[:len("wrapped:")]) != "wrapped:" {
3941
return nil, errors.New("unexpected wrapped dek")
@@ -98,6 +100,36 @@ func TestGCPKMSEnvelopeCipher_RoundTrip(t *testing.T) {
98100
}
99101
}
100102

103+
func TestGCPKMSEnvelopeCipher_DecryptCachesUnwrappedDEK(t *testing.T) {
104+
t.Parallel()
105+
106+
const keyName = "projects/example-project/locations/us-central1/keyRings/example-ring/cryptoKeys/dolthub-auth-staging"
107+
client := &fakeKMSEnvelopeClient{
108+
versionName: keyName + "/cryptoKeyVersions/7",
109+
}
110+
cipher, err := newGCPKMSEnvelopeCipherWithClient(client, keyName)
111+
if err != nil {
112+
t.Fatalf("newGCPKMSEnvelopeCipherWithClient() error = %v", err)
113+
}
114+
115+
encoded, keyVersion, backend, err := cipher.Encrypt(context.Background(), []byte("super-secret-token"))
116+
if err != nil {
117+
t.Fatalf("Encrypt() error = %v", err)
118+
}
119+
for range 2 {
120+
plaintext, err := cipher.Decrypt(context.Background(), encoded, keyVersion, backend)
121+
if err != nil {
122+
t.Fatalf("Decrypt() error = %v", err)
123+
}
124+
if string(plaintext) != "super-secret-token" {
125+
t.Fatalf("plaintext = %q", plaintext)
126+
}
127+
}
128+
if client.decryptCalls != 1 {
129+
t.Fatalf("decryptCalls = %d, want 1", client.decryptCalls)
130+
}
131+
}
132+
101133
func TestNewCredentialCipher_KMSPrimaryWithLegacyLocalFallback(t *testing.T) {
102134
t.Parallel()
103135

internal/dolthubauth/server_test.go

Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -167,6 +167,55 @@ func TestServerRedeemInvalidJSON(t *testing.T) {
167167
}
168168
}
169169

170+
func TestServerRedeemExpiredTokenReturnsCORSJSON(t *testing.T) {
171+
srv, err := NewServer(validConfig(), Dependencies{
172+
Store: expiredRedeemStore{},
173+
KeyManager: fakeCipher{},
174+
})
175+
if err != nil {
176+
t.Fatalf("NewServer() error = %v", err)
177+
}
178+
179+
reqBody, _ := json.Marshal(RedeemConnectTokenRequest{
180+
ConnectToken: "connect-token",
181+
RedeemSecret: "redeem-secret",
182+
APIKey: "secret-token",
183+
Metadata: UserMetadata{
184+
RigHandle: "alice",
185+
Wastelands: []WastelandConfig{{
186+
Upstream: "hop/wl-commons",
187+
ForkOrg: "alice-org",
188+
ForkDB: "wl-commons",
189+
Mode: "pr",
190+
Signing: true,
191+
}},
192+
},
193+
})
194+
req := httptest.NewRequest(http.MethodPost, "/v1/connect-tokens/redeem", bytes.NewReader(reqBody))
195+
req.Header.Set("Origin", "https://app.example")
196+
rec := httptest.NewRecorder()
197+
srv.Handler().ServeHTTP(rec, req)
198+
199+
if rec.Code != http.StatusUnauthorized {
200+
t.Fatalf("status = %d body = %s", rec.Code, rec.Body.String())
201+
}
202+
if got := rec.Header().Get("Access-Control-Allow-Origin"); got != "https://app.example" {
203+
t.Fatalf("Access-Control-Allow-Origin = %q", got)
204+
}
205+
if got := rec.Header().Get("X-Wasteland-Auth-Error-Code"); got != "expired_connect_token" {
206+
t.Fatalf("X-Wasteland-Auth-Error-Code = %q", got)
207+
}
208+
if !strings.Contains(rec.Body.String(), `"error_code":"expired_connect_token"`) {
209+
t.Fatalf("body = %s", rec.Body.String())
210+
}
211+
}
212+
213+
type expiredRedeemStore struct{ fakeStore }
214+
215+
func (f expiredRedeemStore) RedeemConnectToken(context.Context, RedeemInput) (*Connection, error) {
216+
return nil, ErrExpiredConnectToken
217+
}
218+
170219
func TestServerCORSRejectsUnknownOrigin(t *testing.T) {
171220
srv, err := NewServer(validConfig(), Dependencies{
172221
Store: fakeStore{},

internal/dolthubauth/store.go

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -107,7 +107,7 @@ func (s *PostgresStore) RedeemConnectToken(ctx context.Context, input RedeemInpu
107107
if usedAt != nil {
108108
return nil, ErrInvalidConnectToken
109109
}
110-
if input.Now.After(expiresAt) {
110+
if connectTokenExpired(input.Now, expiresAt) {
111111
return nil, ErrExpiredConnectToken
112112
}
113113
if tenantID != input.TenantID || environment != input.Environment {
@@ -558,6 +558,10 @@ func isUniqueViolation(err error) bool {
558558
return errors.As(err, &pgErr) && pgErr.Code == "23505"
559559
}
560560

561+
func connectTokenExpired(now, expiresAt time.Time) bool {
562+
return !now.Before(expiresAt)
563+
}
564+
561565
func (s *PostgresStore) reapExpiredConnectTokens(ctx context.Context, now time.Time) error {
562566
_, err := s.pool.Exec(ctx, `
563567
DELETE FROM connect_tokens

internal/dolthubauth/store_test.go

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@ package dolthubauth
33
import (
44
"errors"
55
"testing"
6+
"time"
67

78
"github.com/jackc/pgx/v5/pgconn"
89
)
@@ -36,3 +37,25 @@ func TestIsUniqueViolation(t *testing.T) {
3637
}
3738
})
3839
}
40+
41+
func TestConnectTokenExpired(t *testing.T) {
42+
expiresAt := time.Date(2026, 4, 25, 5, 30, 0, 0, time.UTC)
43+
44+
tests := []struct {
45+
name string
46+
now time.Time
47+
want bool
48+
}{
49+
{name: "before expiry", now: expiresAt.Add(-time.Nanosecond), want: false},
50+
{name: "at expiry boundary", now: expiresAt, want: true},
51+
{name: "after expiry", now: expiresAt.Add(time.Nanosecond), want: true},
52+
}
53+
54+
for _, tt := range tests {
55+
t.Run(tt.name, func(t *testing.T) {
56+
if got := connectTokenExpired(tt.now, expiresAt); got != tt.want {
57+
t.Fatalf("connectTokenExpired(%s, %s) = %v, want %v", tt.now, expiresAt, got, tt.want)
58+
}
59+
})
60+
}
61+
}

internal/hosted/authservice_resolver_test.go

Lines changed: 34 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,11 +2,15 @@ package hosted
22

33
import (
44
"context"
5+
"io"
56
"net/http"
7+
"strings"
8+
"sync/atomic"
69
"testing"
710
"time"
811

912
"github.com/gastownhall/wasteland/internal/dolthubauth"
13+
"github.com/gastownhall/wasteland/internal/remote"
1014
)
1115

1216
type roundTripFunc func(*http.Request) (*http.Response, error)
@@ -68,3 +72,33 @@ func TestAuthServiceWorkspaceResolver_BuildClientWarmsPendingCache(t *testing.T)
6872
}
6973
cache.Stop()
7074
}
75+
76+
func TestPendingUpstreamCache_DoesNotRefreshWithoutReaders(t *testing.T) {
77+
var calls atomic.Int32
78+
provider := remote.NewDoltHubProviderWithClient(&http.Client{
79+
Transport: roundTripFunc(func(*http.Request) (*http.Response, error) {
80+
calls.Add(1)
81+
return &http.Response{
82+
StatusCode: http.StatusOK,
83+
Header: make(http.Header),
84+
Body: io.NopCloser(strings.NewReader(`{"pulls":[]}`)),
85+
}, nil
86+
}),
87+
})
88+
cache := newPendingUpstreamCache(provider, "hop", "wl-commons", 10*time.Millisecond)
89+
defer cache.Stop()
90+
91+
deadline := time.Now().Add(time.Second)
92+
for calls.Load() == 0 && time.Now().Before(deadline) {
93+
time.Sleep(time.Millisecond)
94+
}
95+
first := calls.Load()
96+
if first == 0 {
97+
t.Fatal("initial pending refresh did not run")
98+
}
99+
100+
time.Sleep(50 * time.Millisecond)
101+
if got := calls.Load(); got != first {
102+
t.Fatalf("pending cache refreshed without readers: calls = %d, want %d", got, first)
103+
}
104+
}

internal/hosted/shared_cache.go

Lines changed: 0 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -58,21 +58,6 @@ func newPendingUpstreamCache(provider *remote.DoltHubProvider, upOrg, upDB strin
5858

5959
c.scheduleRefresh(context.Background())
6060

61-
c.wg.Add(1)
62-
go func() {
63-
defer c.wg.Done()
64-
ticker := time.NewTicker(interval)
65-
defer ticker.Stop()
66-
for {
67-
select {
68-
case <-ticker.C:
69-
c.scheduleRefresh(context.Background())
70-
case <-c.stop:
71-
return
72-
}
73-
}
74-
}()
75-
7661
return c
7762
}
7863

0 commit comments

Comments
 (0)