summaryrefslogtreecommitdiff
path: root/internal/security
diff options
context:
space:
mode:
Diffstat (limited to '')
-rw-r--r--internal/security/credentials.go106
-rw-r--r--internal/security/credentials_test.go28
2 files changed, 100 insertions, 34 deletions
diff --git a/internal/security/credentials.go b/internal/security/credentials.go
index b55fb6b..900261e 100644
--- a/internal/security/credentials.go
+++ b/internal/security/credentials.go
@@ -4,54 +4,114 @@ import (
"crypto/aes"
"crypto/cipher"
"crypto/rand"
+ "crypto/sha256"
"encoding/base64"
"errors"
"fmt"
"io"
+ "strings"
)
type CredentialCipher struct {
+ current credentialKey
+ keys []credentialKey
+}
+
+type credentialKey struct {
+ id []byte
aead cipher.AEAD
}
+var credentialEnvelopeMagic = []byte("AGK1")
+
func NewCredentialCipher(encodedKey string) (*CredentialCipher, error) {
- key, err := base64.StdEncoding.DecodeString(encodedKey)
- if err != nil {
- return nil, fmt.Errorf("decode credential key: %w", err)
- }
- if len(key) != 32 {
- return nil, errors.New("credential key must be a base64-encoded 32-byte key")
- }
- block, err := aes.NewCipher(key)
- if err != nil {
- return nil, fmt.Errorf("create credential cipher: %w", err)
+ return NewCredentialKeyring(encodedKey, nil)
+}
+
+func NewCredentialKeyring(encodedKey string, previous []string) (*CredentialCipher, error) {
+ values := append([]string{encodedKey}, previous...)
+ keys := make([]credentialKey, 0, len(values))
+ seen := map[string]struct{}{}
+ for _, value := range values {
+ value = strings.TrimSpace(value)
+ if value == "" {
+ continue
+ }
+ key, err := base64.StdEncoding.DecodeString(value)
+ if err != nil {
+ return nil, fmt.Errorf("decode credential key: %w", err)
+ }
+ if len(key) != 32 {
+ return nil, errors.New("credential key must be a base64-encoded 32-byte key")
+ }
+ block, err := aes.NewCipher(key)
+ if err != nil {
+ return nil, fmt.Errorf("create credential cipher: %w", err)
+ }
+ aead, err := cipher.NewGCM(block)
+ if err != nil {
+ return nil, fmt.Errorf("create credential AEAD: %w", err)
+ }
+ fingerprint := sha256.Sum256(key)
+ id := base64.RawURLEncoding.EncodeToString(fingerprint[:6])
+ if _, ok := seen[id]; ok {
+ continue
+ }
+ seen[id] = struct{}{}
+ keys = append(keys, credentialKey{id: []byte(id), aead: aead})
}
- aead, err := cipher.NewGCM(block)
- if err != nil {
- return nil, fmt.Errorf("create credential AEAD: %w", err)
+ if len(keys) == 0 {
+ return nil, errors.New("at least one credential key is required")
}
- return &CredentialCipher{aead: aead}, nil
+ return &CredentialCipher{current: keys[0], keys: keys}, nil
}
func (c *CredentialCipher) Encrypt(plaintext string) ([]byte, error) {
if plaintext == "" {
return nil, errors.New("credential cannot be empty")
}
- nonce := make([]byte, c.aead.NonceSize())
+ nonce := make([]byte, c.current.aead.NonceSize())
if _, err := io.ReadFull(rand.Reader, nonce); err != nil {
return nil, fmt.Errorf("generate credential nonce: %w", err)
}
- return c.aead.Seal(nonce, nonce, []byte(plaintext), nil), nil
+ prefix := append(append([]byte{}, credentialEnvelopeMagic...), c.current.id...)
+ sealed := c.current.aead.Seal(nil, nonce, []byte(plaintext), prefix)
+ return append(append(prefix, nonce...), sealed...), nil
}
func (c *CredentialCipher) Decrypt(ciphertext []byte) (string, error) {
- if len(ciphertext) < c.aead.NonceSize() {
- return "", errors.New("credential ciphertext is truncated")
+ if len(ciphertext) >= len(credentialEnvelopeMagic) && string(ciphertext[:len(credentialEnvelopeMagic)]) == string(credentialEnvelopeMagic) {
+ idStart := len(credentialEnvelopeMagic)
+ idEnd := idStart + 8
+ if len(ciphertext) < idEnd {
+ return "", errors.New("credential ciphertext is truncated")
+ }
+ prefix := ciphertext[:idEnd]
+ for _, key := range c.keys {
+ if string(key.id) != string(ciphertext[idStart:idEnd]) {
+ continue
+ }
+ nonceEnd := idEnd + key.aead.NonceSize()
+ if len(ciphertext) < nonceEnd {
+ return "", errors.New("credential ciphertext is truncated")
+ }
+ plaintext, err := key.aead.Open(nil, ciphertext[idEnd:nonceEnd], ciphertext[nonceEnd:], prefix)
+ if err != nil {
+ return "", errors.New("decrypt credential: authentication failed")
+ }
+ return string(plaintext), nil
+ }
+ return "", errors.New("credential key is not present in the configured keyring")
}
- nonce := ciphertext[:c.aead.NonceSize()]
- plaintext, err := c.aead.Open(nil, nonce, ciphertext[c.aead.NonceSize():], nil)
- if err != nil {
- return "", errors.New("decrypt credential: authentication failed")
+ for _, key := range c.keys {
+ if len(ciphertext) < key.aead.NonceSize() {
+ continue
+ }
+ nonce := ciphertext[:key.aead.NonceSize()]
+ plaintext, err := key.aead.Open(nil, nonce, ciphertext[key.aead.NonceSize():], nil)
+ if err == nil {
+ return string(plaintext), nil
+ }
}
- return string(plaintext), nil
+ return "", errors.New("decrypt credential: authentication failed")
}
diff --git a/internal/security/credentials_test.go b/internal/security/credentials_test.go
index 07fa59e..570ec50 100644
--- a/internal/security/credentials_test.go
+++ b/internal/security/credentials_test.go
@@ -2,29 +2,35 @@ package security
import (
"encoding/base64"
- "strings"
"testing"
)
-func TestCredentialCipherRoundTrip(t *testing.T) {
- key := base64.StdEncoding.EncodeToString([]byte(strings.Repeat("k", 32)))
- cipher, err := NewCredentialCipher(key)
+func TestCredentialKeyringDecryptsPreviousAndReencryptsWithPrimary(t *testing.T) {
+ primary := base64.StdEncoding.EncodeToString([]byte("01234567890123456789012345678901"))
+ previous := base64.StdEncoding.EncodeToString([]byte("abcdefghijklmnopqrstuvwxyzabcdef"))
+ oldCipher, err := NewCredentialCipher(previous)
if err != nil {
t.Fatal(err)
}
- ciphertext, err := cipher.Encrypt("upstream-secret")
+ legacy, err := oldCipher.Encrypt("provider-secret")
if err != nil {
t.Fatal(err)
}
- plaintext, err := cipher.Decrypt(ciphertext)
+ keyring, err := NewCredentialKeyring(primary, []string{previous})
if err != nil {
t.Fatal(err)
}
- if plaintext != "upstream-secret" {
- t.Fatalf("unexpected plaintext: %q", plaintext)
+ if got, err := keyring.Decrypt(legacy); err != nil || got != "provider-secret" {
+ t.Fatalf("decrypt previous key: got %q, err %v", got, err)
}
- ciphertext[len(ciphertext)-1] ^= 1
- if _, err := cipher.Decrypt(ciphertext); err == nil {
- t.Fatal("expected authentication failure for modified ciphertext")
+ rotated, err := keyring.Encrypt("provider-secret")
+ if err != nil {
+ t.Fatal(err)
+ }
+ if string(rotated[:4]) != "AGK1" {
+ t.Fatalf("expected versioned credential envelope, got %q", rotated[:4])
+ }
+ if got, err := keyring.Decrypt(rotated); err != nil || got != "provider-secret" {
+ t.Fatalf("decrypt primary key: got %q, err %v", got, err)
}
}