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.
This commit is contained in:
parent
b6638a5067
commit
4dbf3e5fcb
4 changed files with 412 additions and 2 deletions
90
pkg/auth/crypto.go
Normal file
90
pkg/auth/crypto.go
Normal file
|
|
@ -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:<base64>".
|
||||
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)
|
||||
}
|
||||
153
pkg/auth/crypto_test.go
Normal file
153
pkg/auth/crypto_test.go
Normal file
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue