From 4dbf3e5fcb517df509512d7f4a97f19e15dc5e2b Mon Sep 17 00:00:00 2001 From: Paul De Velder Date: Mon, 23 Feb 2026 15:21:02 +0100 Subject: [PATCH] feat(security): add AES-256-GCM secret encryption at rest Encrypt access and refresh tokens in auth.json when PICOCLAW_MASTER_KEY env var is set. Fully backward compatible: without the env var, tokens remain in plaintext. SaveStore clones credentials before encryption to avoid mutating in-memory state. Directory permissions tightened to 0700. --- pkg/auth/crypto.go | 90 +++++++++++++++++++++++ pkg/auth/crypto_test.go | 153 ++++++++++++++++++++++++++++++++++++++++ pkg/auth/store.go | 45 +++++++++++- pkg/auth/store_test.go | 126 +++++++++++++++++++++++++++++++++ 4 files changed, 412 insertions(+), 2 deletions(-) create mode 100644 pkg/auth/crypto.go create mode 100644 pkg/auth/crypto_test.go diff --git a/pkg/auth/crypto.go b/pkg/auth/crypto.go new file mode 100644 index 000000000..356e38258 --- /dev/null +++ b/pkg/auth/crypto.go @@ -0,0 +1,90 @@ +package auth + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "fmt" + "io" + "os" + "strings" +) + +const encryptedPrefix = "enc:v1:" + +// DeriveKey derives a 32-byte AES-256 key from a master secret using SHA-256. +func DeriveKey(masterSecret string) []byte { + h := sha256.Sum256([]byte(masterSecret)) + return h[:] +} + +// GetMasterKey reads the master key from the PICOCLAW_MASTER_KEY environment variable. +func GetMasterKey() string { + return os.Getenv("PICOCLAW_MASTER_KEY") +} + +// Encrypt encrypts plaintext using AES-256-GCM and returns "enc:v1:". +func Encrypt(plaintext string, key []byte) (string, error) { + block, err := aes.NewCipher(key) + if err != nil { + return "", fmt.Errorf("create cipher: %w", err) + } + + gcm, err := cipher.NewGCM(block) + if err != nil { + return "", fmt.Errorf("create GCM: %w", err) + } + + nonce := make([]byte, gcm.NonceSize()) + if _, err := io.ReadFull(rand.Reader, nonce); err != nil { + return "", fmt.Errorf("generate nonce: %w", err) + } + + ciphertext := gcm.Seal(nonce, nonce, []byte(plaintext), nil) + encoded := base64.StdEncoding.EncodeToString(ciphertext) + return encryptedPrefix + encoded, nil +} + +// Decrypt decrypts a value encrypted by Encrypt. If the value does not have the +// encrypted prefix, it is returned unchanged (plaintext passthrough). +func Decrypt(value string, key []byte) (string, error) { + if !IsEncrypted(value) { + return value, nil + } + + encoded := strings.TrimPrefix(value, encryptedPrefix) + ciphertext, err := base64.StdEncoding.DecodeString(encoded) + if err != nil { + return "", fmt.Errorf("decode base64: %w", err) + } + + block, err := aes.NewCipher(key) + if err != nil { + return "", fmt.Errorf("create cipher: %w", err) + } + + gcm, err := cipher.NewGCM(block) + if err != nil { + return "", fmt.Errorf("create GCM: %w", err) + } + + nonceSize := gcm.NonceSize() + if len(ciphertext) < nonceSize { + return "", fmt.Errorf("ciphertext too short") + } + + nonce, ciphertext := ciphertext[:nonceSize], ciphertext[nonceSize:] + plaintext, err := gcm.Open(nil, nonce, ciphertext, nil) + if err != nil { + return "", fmt.Errorf("decrypt: %w", err) + } + + return string(plaintext), nil +} + +// IsEncrypted returns true if the value has the encrypted prefix. +func IsEncrypted(value string) bool { + return strings.HasPrefix(value, encryptedPrefix) +} diff --git a/pkg/auth/crypto_test.go b/pkg/auth/crypto_test.go new file mode 100644 index 000000000..b7484ba99 --- /dev/null +++ b/pkg/auth/crypto_test.go @@ -0,0 +1,153 @@ +package auth + +import ( + "strings" + "testing" +) + +func TestDeriveKey_Deterministic(t *testing.T) { + key1 := DeriveKey("my-secret") + key2 := DeriveKey("my-secret") + if string(key1) != string(key2) { + t.Error("DeriveKey should be deterministic for same input") + } +} + +func TestDeriveKey_DifferentInputs(t *testing.T) { + key1 := DeriveKey("secret-a") + key2 := DeriveKey("secret-b") + if string(key1) == string(key2) { + t.Error("DeriveKey should produce different keys for different inputs") + } +} + +func TestDeriveKey_Length(t *testing.T) { + key := DeriveKey("test") + if len(key) != 32 { + t.Errorf("DeriveKey should produce 32-byte key, got %d", len(key)) + } +} + +func TestEncryptDecrypt_RoundTrip(t *testing.T) { + key := DeriveKey("master-key") + plaintext := "my-secret-token-12345" + + encrypted, err := Encrypt(plaintext, key) + if err != nil { + t.Fatalf("Encrypt failed: %v", err) + } + + if encrypted == plaintext { + t.Error("Encrypted value should differ from plaintext") + } + + if !IsEncrypted(encrypted) { + t.Error("Encrypted value should have enc:v1: prefix") + } + + decrypted, err := Decrypt(encrypted, key) + if err != nil { + t.Fatalf("Decrypt failed: %v", err) + } + + if decrypted != plaintext { + t.Errorf("Decrypt mismatch: got %q, want %q", decrypted, plaintext) + } +} + +func TestDecrypt_PlaintextPassthrough(t *testing.T) { + key := DeriveKey("any-key") + plaintext := "not-encrypted-value" + + result, err := Decrypt(plaintext, key) + if err != nil { + t.Fatalf("Decrypt of plaintext should not error: %v", err) + } + if result != plaintext { + t.Errorf("Plaintext passthrough failed: got %q, want %q", result, plaintext) + } +} + +func TestDecrypt_WrongKey(t *testing.T) { + key1 := DeriveKey("correct-key") + key2 := DeriveKey("wrong-key") + + encrypted, err := Encrypt("secret", key1) + if err != nil { + t.Fatalf("Encrypt failed: %v", err) + } + + _, err = Decrypt(encrypted, key2) + if err == nil { + t.Error("Decrypt with wrong key should fail") + } +} + +func TestIsEncrypted(t *testing.T) { + tests := []struct { + value string + expected bool + }{ + {"enc:v1:abcdef", true}, + {"enc:v1:", true}, + {"plaintext", false}, + {"enc:v2:something", false}, + {"", false}, + } + for _, tc := range tests { + if got := IsEncrypted(tc.value); got != tc.expected { + t.Errorf("IsEncrypted(%q) = %v, want %v", tc.value, got, tc.expected) + } + } +} + +func TestEncrypt_UniqueNonces(t *testing.T) { + key := DeriveKey("key") + plaintext := "same-value" + + enc1, _ := Encrypt(plaintext, key) + enc2, _ := Encrypt(plaintext, key) + + if enc1 == enc2 { + t.Error("Two encryptions of same plaintext should differ (unique nonces)") + } +} + +func TestEncrypt_EmptyString(t *testing.T) { + key := DeriveKey("key") + + encrypted, err := Encrypt("", key) + if err != nil { + t.Fatalf("Encrypt empty string failed: %v", err) + } + + decrypted, err := Decrypt(encrypted, key) + if err != nil { + t.Fatalf("Decrypt empty string failed: %v", err) + } + if decrypted != "" { + t.Errorf("Expected empty string, got %q", decrypted) + } +} + +func TestGetMasterKey_EnvVar(t *testing.T) { + t.Setenv("PICOCLAW_MASTER_KEY", "test-master") + if got := GetMasterKey(); got != "test-master" { + t.Errorf("GetMasterKey() = %q, want %q", got, "test-master") + } +} + +func TestGetMasterKey_Unset(t *testing.T) { + t.Setenv("PICOCLAW_MASTER_KEY", "") + if got := GetMasterKey(); got != "" { + t.Errorf("GetMasterKey() should be empty when env unset, got %q", got) + } +} + +func TestEncryptedPrefix(t *testing.T) { + key := DeriveKey("key") + encrypted, _ := Encrypt("test", key) + if !strings.HasPrefix(encrypted, "enc:v1:") { + t.Errorf("Encrypted value should start with 'enc:v1:', got %q", encrypted) + } +} diff --git a/pkg/auth/store.go b/pkg/auth/store.go index 64708421b..bb3be6376 100644 --- a/pkg/auth/store.go +++ b/pkg/auth/store.go @@ -58,17 +58,58 @@ func LoadStore() (*AuthStore, error) { if store.Credentials == nil { store.Credentials = make(map[string]*AuthCredential) } + + // Decrypt tokens if master key is set + if mk := GetMasterKey(); mk != "" { + key := DeriveKey(mk) + for _, cred := range store.Credentials { + if IsEncrypted(cred.AccessToken) { + if dec, err := Decrypt(cred.AccessToken, key); err == nil { + cred.AccessToken = dec + } + } + if IsEncrypted(cred.RefreshToken) { + if dec, err := Decrypt(cred.RefreshToken, key); err == nil { + cred.RefreshToken = dec + } + } + } + } + return &store, nil } func SaveStore(store *AuthStore) error { path := authFilePath() dir := filepath.Dir(path) - if err := os.MkdirAll(dir, 0o755); err != nil { + if err := os.MkdirAll(dir, 0o700); err != nil { return err } - data, err := json.MarshalIndent(store, "", " ") + // Clone and encrypt tokens if master key is set + toSave := &AuthStore{Credentials: make(map[string]*AuthCredential, len(store.Credentials))} + for k, cred := range store.Credentials { + clone := *cred + toSave.Credentials[k] = &clone + } + + if mk := GetMasterKey(); mk != "" { + key := DeriveKey(mk) + for _, cred := range toSave.Credentials { + if cred.AccessToken != "" && !IsEncrypted(cred.AccessToken) { + if enc, err := Encrypt(cred.AccessToken, key); err == nil { + cred.AccessToken = enc + } + } + if cred.RefreshToken != "" && !IsEncrypted(cred.RefreshToken) { + if enc, err := Encrypt(cred.RefreshToken, key); err == nil { + cred.RefreshToken = enc + } + } + } + } + + data, err := json.MarshalIndent(toSave, "", " ") if err != nil { return err } diff --git a/pkg/auth/store_test.go b/pkg/auth/store_test.go index f6793cfce..bbd991fe2 100644 --- a/pkg/auth/store_test.go +++ b/pkg/auth/store_test.go @@ -1,6 +1,7 @@ package auth import ( + "encoding/json" "os" "path/filepath" "testing" @@ -187,3 +188,128 @@ func TestLoadStoreEmpty(t *testing.T) { t.Errorf("expected empty credentials, got %d", len(store.Credentials)) } } + +// --- Encryption integration tests --- + +func TestStoreRoundtrip_WithEncryption(t *testing.T) { + tmpDir := t.TempDir() + t.Setenv("HOME", tmpDir) + t.Setenv("USERPROFILE", tmpDir) + t.Setenv("PICOCLAW_MASTER_KEY", "test-master-key") + + cred := &AuthCredential{ + AccessToken: "secret-access-token", + RefreshToken: "secret-refresh-token", + Provider: "openai", + AuthMethod: "oauth", + } + + if err := SetCredential("openai", cred); err != nil { + t.Fatalf("SetCredential() error: %v", err) + } + + // Verify on-disk tokens are encrypted + path := filepath.Join(tmpDir, ".picoclaw", "auth.json") + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("ReadFile() error: %v", err) + } + + var raw struct { + Credentials map[string]json.RawMessage `json:"credentials"` + } + if err := json.Unmarshal(data, &raw); err != nil { + t.Fatalf("Unmarshal raw: %v", err) + } + + var rawCred struct { + AccessToken string `json:"access_token"` + RefreshToken string `json:"refresh_token"` + } + if err := json.Unmarshal(raw.Credentials["openai"], &rawCred); err != nil { + t.Fatalf("Unmarshal cred: %v", err) + } + + if !IsEncrypted(rawCred.AccessToken) { + t.Errorf("On-disk access_token should be encrypted, got: %s", rawCred.AccessToken) + } + if !IsEncrypted(rawCred.RefreshToken) { + t.Errorf("On-disk refresh_token should be encrypted, got: %s", rawCred.RefreshToken) + } + + // Verify LoadStore decrypts correctly + loaded, err := GetCredential("openai") + if err != nil { + t.Fatalf("GetCredential() error: %v", err) + } + if loaded.AccessToken != "secret-access-token" { + t.Errorf("AccessToken = %q, want %q", loaded.AccessToken, "secret-access-token") + } + if loaded.RefreshToken != "secret-refresh-token" { + t.Errorf("RefreshToken = %q, want %q", loaded.RefreshToken, "secret-refresh-token") + } +} + +func TestStoreRoundtrip_WithoutMasterKey(t *testing.T) { + tmpDir := t.TempDir() + t.Setenv("HOME", tmpDir) + t.Setenv("USERPROFILE", tmpDir) + t.Setenv("PICOCLAW_MASTER_KEY", "") + + cred := &AuthCredential{ + AccessToken: "plaintext-token", + Provider: "openai", + AuthMethod: "oauth", + } + + if err := SetCredential("openai", cred); err != nil { + t.Fatalf("SetCredential() error: %v", err) + } + + // Verify on-disk token is plaintext (no encryption without master key) + path := filepath.Join(tmpDir, ".picoclaw", "auth.json") + data, err := os.ReadFile(path) + if err != nil { + t.Fatalf("ReadFile() error: %v", err) + } + + var raw struct { + Credentials map[string]json.RawMessage `json:"credentials"` + } + json.Unmarshal(data, &raw) + + var rawCred struct { + AccessToken string `json:"access_token"` + } + json.Unmarshal(raw.Credentials["openai"], &rawCred) + + if IsEncrypted(rawCred.AccessToken) { + t.Error("Without master key, token should be plaintext on disk") + } + if rawCred.AccessToken != "plaintext-token" { + t.Errorf("AccessToken = %q, want %q", rawCred.AccessToken, "plaintext-token") + } +} + +func TestStore_DirPermissions(t *testing.T) { + tmpDir := t.TempDir() + t.Setenv("HOME", tmpDir) + t.Setenv("USERPROFILE", tmpDir) + t.Setenv("PICOCLAW_MASTER_KEY", "") + + cred := &AuthCredential{AccessToken: "x", Provider: "test", AuthMethod: "token"} + if err := SetCredential("test", cred); err != nil { + t.Fatalf("SetCredential() error: %v", err) + } + + dir := filepath.Join(tmpDir, ".picoclaw") + info, err := os.Stat(dir) + if err != nil { + t.Fatalf("Stat dir: %v", err) + } + perm := info.Mode().Perm() + // On Windows, permissions may differ; check that it's at least created + if perm&0o700 != 0o700 { + t.Logf("Directory permissions = %o (platform may differ)", perm) + } +}