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 {
|
if store.Credentials == nil {
|
||||||
store.Credentials = make(map[string]*AuthCredential)
|
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
|
return &store, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func SaveStore(store *AuthStore) error {
|
func SaveStore(store *AuthStore) error {
|
||||||
path := authFilePath()
|
path := authFilePath()
|
||||||
dir := filepath.Dir(path)
|
dir := filepath.Dir(path)
|
||||||
if err := os.MkdirAll(dir, 0o755); err != nil {
|
if err := os.MkdirAll(dir, 0o700); err != nil {
|
||||||
return err
|
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 {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,6 +1,7 @@
|
||||||
package auth
|
package auth
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"encoding/json"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
@ -187,3 +188,128 @@ func TestLoadStoreEmpty(t *testing.T) {
|
||||||
t.Errorf("expected empty credentials, got %d", len(store.Credentials))
|
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