From 5986d3bb494450f2a0b46751452abe8327a248e1 Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Sat, 15 Aug 2026 20:45:47 +0530 Subject: [PATCH 1/3] feat(storage): at-rest encryption for node content via key providers migration v5's columns are now wired to a real key provider: - KeyProvider interface: versioned master key material; EnvKeyProvider reads it from a process environment variable (no secret on disk) - NodeCipher: AES-256-GCM with per-version keys derived via HKDF-SHA256; ciphertexts carry 'yaad.aes256gcm.v{n}.' prefix, random nonce, and the key version as authenticated data (relabeling/tampering fails reads) - Store.EnableEncryption(opts in after NewStore) threads the cipher through every node read/write path, txStore included, and bookkeeps nodes.encrypted / encryption_key_version - node_versions history encrypts alongside live content - legacy plaintext rows stay readable after enabling encryption and are progressively re-encrypted on update - SearchNodes falls back to an in-memory token scan over decrypted content when the cipher is active (FTS5 indexes ciphertext) - default behaviour unchanged: no provider = plaintext, columns false/0 - go directive -> 1.26.6 (stdlib CVEs per govulncheck) tests: round trip, tamper + wrong-key rejection, env provider, legacy plaintext boundary, tx path, FTS fallback; full suite + race green --- go.mod | 5 +- go.sum | 14 +- storage/crypto.go | 223 +++++++++++++++++++ storage/crypto_test.go | 466 ++++++++++++++++++++++++++++++++++++++++ storage/errors.go | 5 + storage/node_crypto.go | 94 ++++++++ storage/prefix.go | 2 +- storage/sqlite.go | 44 +++- storage/sqlite_edges.go | 6 +- storage/sqlite_nodes.go | 155 ++++++++++--- storage/sqlite_tx.go | 42 ++-- 11 files changed, 993 insertions(+), 63 deletions(-) create mode 100644 storage/crypto.go create mode 100644 storage/crypto_test.go create mode 100644 storage/node_crypto.go diff --git a/go.mod b/go.mod index 012f4a9..3ee17fe 100644 --- a/go.mod +++ b/go.mod @@ -1,6 +1,6 @@ module github.com/GrayCodeAI/yaad -go 1.26.5 +go 1.26.6 require ( github.com/BurntSushi/toml v1.6.0 @@ -13,8 +13,9 @@ require ( go.opentelemetry.io/otel v1.44.0 go.opentelemetry.io/otel/metric v1.44.0 go.opentelemetry.io/otel/sdk/metric v1.44.0 + golang.org/x/crypto v0.55.0 golang.org/x/sys v0.47.0 - golang.org/x/text v0.40.0 + golang.org/x/text v0.41.0 google.golang.org/grpc v1.82.1 google.golang.org/protobuf v1.36.11 modernc.org/sqlite v1.51.0 diff --git a/go.sum b/go.sum index b348993..5da43cc 100644 --- a/go.sum +++ b/go.sum @@ -81,18 +81,20 @@ go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ= go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ= -golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ= -golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= +golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M= +golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis= +golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk= +golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40= golang.org/x/net v0.57.0 h1:K5+3DljvIuDG9/Jv9rvyMywYNFCQ9RSUY6OOTTkT+tE= golang.org/x/net v0.57.0/go.mod h1:KpXc8iv+r3XplLAG/f7Jsf9RPszJzdR0f58q9vGOuEU= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= -golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= -golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= -golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q= -golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE= +golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk= gonum.org/v1/gonum v0.17.0 h1:VbpOemQlsSMrYmn7T2OUvQ4dqxQXU+ouZFQsZOx50z4= gonum.org/v1/gonum v0.17.0/go.mod h1:El3tOrEuMpv2UdMrbNlKEh9vd86bmQ6vqIcDwxEOc1E= google.golang.org/genproto/googleapis/rpc v0.0.0-20260414002931-afd174a4e478 h1:RmoJA1ujG+/lRGNfUnOMfhCy5EipVMyvUE+KNbPbTlw= diff --git a/storage/crypto.go b/storage/crypto.go new file mode 100644 index 0000000..1242814 --- /dev/null +++ b/storage/crypto.go @@ -0,0 +1,223 @@ +package storage + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "crypto/sha256" + "encoding/base64" + "fmt" + "io" + "os" + "strconv" + "strings" + "sync" + + "golang.org/x/crypto/hkdf" +) + +// CipherScheme identifies the at-rest encryption format used for node +// content. The scheme is embedded in every ciphertext so future formats can +// coexist with current ones on reads. +const CipherScheme = "yaad.aes256gcm.v" + +// CiphertextPrefix is the full marker stored at the head of encrypted node +// content. Content without this prefix is treated as legacy plaintext. +// +// The prefix is deliberately namespaced and self-describing: FTS5 and LIKE +// queries against encrypted rows simply never match plaintext queries, and a +// stored value can always be told apart from content a user typed. +const CiphertextPrefix = CipherScheme + +// hkdfSaltDomain binds derived keys to this column domain. Rotating keys or +// adding encrypted columns later should change this domain. +const hkdfSaltDomain = "yaad node content at rest" + +// KeyProvider supplies master key material for at-rest encryption. It is the +// security boundary between the storage layer and wherever the real secret +// lives — environment, keychain, KMS. Implementations return raw key bytes +// per version; CurrentVersion names the version used for new writes and +// older versions remain available so rotation can re-encrypt progressively. +type KeyProvider interface { + // KeyBytes returns raw master key material for the requested version. + // Unknown versions return an error; empty material is rejected. + KeyBytes(version int) ([]byte, error) + // CurrentVersion is the key version used to encrypt new data. + CurrentVersion() int +} + +// EnvKeyProvider is a KeyProvider backed by an environment variable. It is +// the wired "real key provider" for yaad's encryption columns: the secret +// never touches disk, only the caller's process environment. Only version 0 +// is supported — hosts that need rotation should layer their own provider. +type EnvKeyProvider struct { + envVar string +} + +// NewEnvKeyProvider returns a provider reading the master key from envVar. +// The variable is read lazily on each KeyBytes call so tests can set it via +// t.Setenv after construction. +func NewEnvKeyProvider(envVar string) *EnvKeyProvider { + return &EnvKeyProvider{envVar: envVar} +} + +// CurrentVersion implements KeyProvider. +func (p *EnvKeyProvider) CurrentVersion() int { return 0 } + +// KeyBytes implements KeyProvider. +func (p *EnvKeyProvider) KeyBytes(version int) ([]byte, error) { + if version != 0 { + return nil, fmt.Errorf("%w: %d", ErrUnknownKeyVersion, version) + } + raw := strings.TrimSpace(os.Getenv(p.envVar)) + if raw == "" { + return nil, fmt.Errorf("%w: %s is unset or empty", ErrKeyMaterial, p.envVar) + } + return []byte(raw), nil +} + +// NodeCipher encrypts and decrypts single string values with +// AES-256-GCM. Per-version encryption keys are derived from provider master +// material via HKDF-SHA256 (deterministic per version, never stored); each +// encrypted value carries a random 12-byte nonce and authenticates the key +// version as associated data, so ciphertexts cannot be relabeled across +// versions. Encrypted output is ASCII (prefix + base64url) so it round-trips +// through TEXT columns, backups, and JSON without encoding changes. +// +// A NodeCipher is safe for concurrent use. +type NodeCipher struct { + provider KeyProvider + current int // key version for new writes, captured at construction + + mu sync.RWMutex + derived map[int]cipher.AEAD // lazily derived per version +} + +// NewNodeCipher validates the provider eagerly: the current key version is +// derived now so misconfiguration (missing env var, bad material) fails at +// setup instead of on the first write. +func NewNodeCipher(p KeyProvider) (*NodeCipher, error) { + if p == nil { + return nil, fmt.Errorf("nil key provider") + } + c := &NodeCipher{ + provider: p, + current: p.CurrentVersion(), + derived: make(map[int]cipher.AEAD), + } + if _, err := c.aeadFor(c.current); err != nil { + return nil, err + } + return c, nil +} + +// KeyVersion returns the key version used for new encrypted values. +func (c *NodeCipher) KeyVersion() int { return c.current } + +// Encrypt seals plaintext with the current key version. Empty strings pass +// through unchanged so nullable columns stay empty rather than becoming a +// decryptable "empty" blob. +func (c *NodeCipher) Encrypt(plaintext string) (string, error) { + if plaintext == "" { + return "", nil + } + aead, err := c.aeadFor(c.current) + if err != nil { + return "", err + } + nonce := make([]byte, aead.NonceSize()) + if _, err := rand.Read(nonce); err != nil { + return "", fmt.Errorf("generate nonce: %w", err) + } + sealed := aead.Seal(nil, nonce, []byte(plaintext), versionAAD(c.current)) + + payload := make([]byte, 0, len(nonce)+len(sealed)) + payload = append(payload, nonce...) + payload = append(payload, sealed...) + return CipherScheme + strconv.Itoa(c.current) + "." + base64.RawURLEncoding.EncodeToString(payload), nil +} + +// Decrypt opens a ciphertext produced by Encrypt. Values without the cipher +// prefix are returned unchanged — this is how a store transparently reads +// legacy plaintext rows written before encryption was enabled. +func (c *NodeCipher) Decrypt(stored string) (string, error) { + if !IsEncryptedValue(stored) { + return stored, nil + } + rest := stored[len(CipherScheme):] + dot := strings.IndexByte(rest, '.') + if dot < 1 { + return "", fmt.Errorf("%w: missing version separator", ErrMalformedCiphertext) + } + version, err := strconv.Atoi(rest[:dot]) + if err != nil { + return "", fmt.Errorf("%w: bad key version %q", ErrMalformedCiphertext, rest[:dot]) + } + payload, err := base64.RawURLEncoding.DecodeString(rest[dot+1:]) + if err != nil { + return "", fmt.Errorf("%w: bad base64 payload", ErrMalformedCiphertext) + } + aead, err := c.aeadFor(version) + if err != nil { + return "", err + } + ns := aead.NonceSize() + if len(payload) < ns+aead.Overhead() { + return "", fmt.Errorf("%w: payload shorter than nonce+tag", ErrMalformedCiphertext) + } + plaintext, err := aead.Open(nil, payload[:ns], payload[ns:], versionAAD(version)) + if err != nil { + return "", fmt.Errorf("%w", ErrDecryptionFailed) + } + return string(plaintext), nil +} + +// IsEncryptedValue reports whether a stored string carries the cipher prefix. +func IsEncryptedValue(s string) bool { + return strings.HasPrefix(s, CiphertextPrefix) +} + +// aeadFor returns (and caches) the AES-256-GCM AEAD derived from version's +// master material. Double-checked under the mutex so hot reads avoid the +// derivation path entirely. +func (c *NodeCipher) aeadFor(version int) (cipher.AEAD, error) { + c.mu.RLock() + a := c.derived[version] + c.mu.RUnlock() + if a != nil { + return a, nil + } + + c.mu.Lock() + defer c.mu.Unlock() + if a := c.derived[version]; a != nil { + return a, nil + } + master, err := c.provider.KeyBytes(version) + if err != nil { + return nil, err + } + if len(master) == 0 { + return nil, fmt.Errorf("%w: provider returned empty material for version %d", ErrKeyMaterial, version) + } + hk := hkdf.New(sha256.New, master, []byte(hkdfSaltDomain), []byte("aes256gcm key v"+strconv.Itoa(version))) + key := make([]byte, 32) + if _, err := io.ReadFull(hk, key); err != nil { + return nil, fmt.Errorf("derive key: %w", err) + } + block, err := aes.NewCipher(key) + if err != nil { + return nil, fmt.Errorf("aes cipher: %w", err) + } + aead, err := cipher.NewGCM(block) + if err != nil { + return nil, fmt.Errorf("gcm: %w", err) + } + c.derived[version] = aead + return aead, nil +} + +// versionAAD binds each ciphertext to the key version that sealed it. +func versionAAD(version int) []byte { + return []byte("yaad-node-v" + strconv.Itoa(version)) +} diff --git a/storage/crypto_test.go b/storage/crypto_test.go new file mode 100644 index 0000000..d460dc9 --- /dev/null +++ b/storage/crypto_test.go @@ -0,0 +1,466 @@ +package storage + +import ( + "context" + "database/sql" + "errors" + "fmt" + "strings" + "testing" +) + +type constKeyProvider struct { + key []byte + version int +} + +func (p *constKeyProvider) KeyBytes(version int) ([]byte, error) { + if version != p.version { + return nil, ErrUnknownKeyVersion + } + return p.key, nil +} + +func (p *constKeyProvider) CurrentVersion() int { return p.version } + +func newTestCipher(t *testing.T, v int) *NodeCipher { + t.Helper() + c, err := NewNodeCipher(&constKeyProvider{key: []byte("unit-test-master-key-material-32b"), version: v}) + if err != nil { + t.Fatal(err) + } + return c +} + +func TestNodeCipherRoundTrip(t *testing.T) { + c := newTestCipher(t, 0) + + plain := "convention: always run make lint before opening a PR" + enc, err := c.Encrypt(plain) + if err != nil { + t.Fatal(err) + } + if !IsEncryptedValue(enc) { + t.Fatalf("ciphertext missing prefix: %q", enc) + } + if strings.Contains(enc, plain) { + t.Fatal("plaintext leaked into ciphertext") + } + enc2, err := c.Encrypt(plain) + if err != nil { + t.Fatal(err) + } + if enc == enc2 { + t.Fatal("deterministic ciphertext: nonce must randomize") + } + dec, err := c.Decrypt(enc) + if err != nil { + t.Fatal(err) + } + if dec != plain { + t.Fatalf("round trip: got %q", dec) + } +} + +func TestNodeCipherPassthrough(t *testing.T) { + c := newTestCipher(t, 0) + + got, err := c.Decrypt("legacy plaintext content") + if err != nil { + t.Fatal(err) + } + if got != "legacy plaintext content" { + t.Fatalf("plaintext passthrough failed: %q", got) + } + + if enc, err := c.Encrypt(""); err != nil || enc != "" { + t.Fatalf("empty plaintext should pass through: %q, %v", enc, err) + } +} + +func TestNodeCipherTamper(t *testing.T) { + c := newTestCipher(t, 0) + enc, err := c.Encrypt("sensitive") + if err != nil { + t.Fatal(err) + } + // Flip one character in the base64 payload (GCM auth must reject it). + i := len(enc) - 5 + tampered := enc[:i] + "A" + enc[i+1:] + if _, err := c.Decrypt(tampered); !errors.Is(err, ErrDecryptionFailed) { + t.Fatalf("expected ErrDecryptionFailed, got %v", err) + } +} + +func TestNodeCipherVersionBinding(t *testing.T) { + c0 := newTestCipher(t, 0) + // c0b uses DIFFERENT key material for the same version 0 — a wrong key + // must fail authentication, and relabeling the version must too. + c0b, err := NewNodeCipher(&constKeyProvider{key: []byte("different-key-material-32-bytes!!"), version: 0}) + if err != nil { + t.Fatal(err) + } + + enc0, err := c0.Encrypt("secret v0") + if err != nil { + t.Fatal(err) + } + if _, err := c0b.Decrypt(enc0); !errors.Is(err, ErrDecryptionFailed) { + t.Fatalf("wrong-key decryption must fail, got %v", err) + } +} + +func TestNodeCipherRelabeledVersionRejected(t *testing.T) { + c0 := newTestCipher(t, 0) + enc, err := c0.Encrypt("bind the version") + if err != nil { + t.Fatal(err) + } + relabeled := strings.Replace(enc, CipherScheme+"0.", CipherScheme+"1.", 1) + if _, err := c0.Decrypt(relabeled); !errors.Is(err, ErrUnknownKeyVersion) { + t.Fatalf("relabeled ciphertext must be rejected (unknown version), got %v", err) + } +} + +func TestNewNodeCipherEagerValidation(t *testing.T) { + t.Setenv("YAAD_TEST_MISSING_KEY", "") + if _, err := NewNodeCipher(NewEnvKeyProvider("YAAD_TEST_MISSING_KEY")); !errors.Is(err, ErrKeyMaterial) { + t.Fatalf("expected ErrKeyMaterial for unset env var, got %v", err) + } +} + +func TestEnvKeyProvider(t *testing.T) { + t.Setenv("YAAD_TEST_KEY", "env-provided-key-material") + p := NewEnvKeyProvider("YAAD_TEST_KEY") + if p.CurrentVersion() != 0 { + t.Fatalf("env provider version: %d", p.CurrentVersion()) + } + b, err := p.KeyBytes(0) + if err != nil || string(b) != "env-provided-key-material" { + t.Fatalf("KeyBytes(0) = %q, %v", b, err) + } + if _, err := p.KeyBytes(1); !errors.Is(err, ErrUnknownKeyVersion) { + t.Fatalf("KeyBytes(1) should be unknown version, got %v", err) + } + + c, err := NewNodeCipher(p) + if err != nil { + t.Fatal(err) + } + enc, err := c.Encrypt("from env") + if err != nil { + t.Fatal(err) + } + dec, err := c.Decrypt(enc) + if err != nil || dec != "from env" { + t.Fatalf("env-key round trip: %q, %v", dec, err) + } +} + +// TestStoreEncryptionRoundTrip covers the full storage path: writes encrypt, +// reads decrypt, and the bookkeeping columns carry the right flags. +func TestStoreEncryptionRoundTrip(t *testing.T) { + s, cleanup := setupStore(t) + defer cleanup() + if err := s.EnableEncryption(&constKeyProvider{key: []byte("store-round-trip-key-material!!"), version: 0}); err != nil { + t.Fatal(err) + } + ctx := context.Background() + + n := &Node{ + ID: "enc-1", Type: "convention", Content: "encrypt me at rest", + Summary: "summary too", ContentHash: "h-enc-1", Scope: "project", Project: "test", + Key: "enc-key-1", + } + if err := s.CreateNode(ctx, n); err != nil { + t.Fatal(err) + } + + got, err := s.GetNode(ctx, "enc-1") + if err != nil { + t.Fatal(err) + } + if got.Content != "encrypt me at rest" || got.Summary != "summary too" { + t.Fatalf("API returned wrong plaintext: %q / %q", got.Content, got.Summary) + } + + // At rest: raw content must be ciphertext with the bookkeeping set. + var rawContent, rawSummary string + var encrypted bool + var keyVersion int + err = s.DB().QueryRowContext(ctx, + `SELECT content, summary, encrypted, encryption_key_version FROM nodes WHERE id=?`, "enc-1"). + Scan(&rawContent, &rawSummary, &encrypted, &keyVersion) + if err != nil { + t.Fatal(err) + } + if !IsEncryptedValue(rawContent) || !IsEncryptedValue(rawSummary) { + t.Fatalf("at-rest values not encrypted: %q / %q", rawContent, rawSummary) + } + if strings.Contains(rawContent, "encrypt me at rest") { + t.Fatal("plaintext visible at rest") + } + if !encrypted || keyVersion != 0 { + t.Fatalf("bookkeeping wrong: encrypted=%v keyVersion=%d", encrypted, keyVersion) + } + + // Other readers decrypt too. + list, err := s.ListNodes(ctx, NodeFilter{Type: "convention"}) + if err != nil || len(list) != 1 || list[0].Content != "encrypt me at rest" { + t.Fatalf("ListNodes decrypt failed: %v, %+v", err, list) + } + byKey, err := s.GetNodeByKey(ctx, "enc-key-1", "test") + if err != nil || byKey == nil || byKey.Content != "encrypt me at rest" { + t.Fatalf("GetNodeByKey decrypt failed: %v, %+v", err, byKey) + } + batch, err := s.GetNodesBatch(ctx, []string{"enc-1"}) + if err != nil || len(batch) != 1 || batch[0].Content != "encrypt me at rest" { + t.Fatalf("GetNodesBatch decrypt failed: %v", err) + } + + // Updates re-encrypt. + got.Content = "updated plaintext" + if err := s.UpdateNode(ctx, got); err != nil { + t.Fatal(err) + } + var raw2 string + if err := s.DB().QueryRowContext(ctx, `SELECT content FROM nodes WHERE id=?`, "enc-1").Scan(&raw2); err != nil { + t.Fatal(err) + } + if strings.Contains(raw2, "updated plaintext") { + t.Fatal("update stored plaintext at rest") + } + got2, err := s.GetNode(ctx, "enc-1") + if err != nil || got2.Content != "updated plaintext" { + t.Fatalf("read after update: %v, %q", err, got2.Content) + } + + if err := s.UpdateNodeContent(ctx, "enc-1", "content-only update"); err != nil { + t.Fatal(err) + } + got3, err := s.GetNode(ctx, "enc-1") + if err != nil || got3.Content != "content-only update" { + t.Fatalf("read after UpdateNodeContent: %v, %q", err, got3.Content) + } + + // Version history encrypts and decrypts. + if err := s.SaveVersion(ctx, "enc-1", "historical plaintext", "test", "v1"); err != nil { + t.Fatal(err) + } + vers, err := s.GetVersions(ctx, "enc-1") + if err != nil || len(vers) != 1 || vers[0].Content != "historical plaintext" { + t.Fatalf("GetVersions decrypt failed: %v", err) + } + var rawVer string + if err := s.DB().QueryRowContext(ctx, + `SELECT content FROM node_versions WHERE node_id=? AND version=1`, "enc-1").Scan(&rawVer); err != nil { + t.Fatal(err) + } + if !IsEncryptedValue(rawVer) { + t.Fatalf("version history not encrypted at rest: %q", rawVer) + } +} + +// TestStoreEncryptionPlaintextLegacy covers the migration boundary: rows +// written before EnableEncryption remain readable afterwards, and a rewrite +// upgrades them to ciphertext. +func TestStoreEncryptionPlaintextLegacy(t *testing.T) { + s, cleanup := setupStore(t) + defer cleanup() + ctx := context.Background() + + legacy := &Node{ + ID: "legacy-1", Type: "decision", Content: "written in plaintext", + ContentHash: "h-legacy", Scope: "project", Project: "test", + } + if err := s.CreateNode(ctx, legacy); err != nil { + t.Fatal(err) + } + + if err := s.EnableEncryption(&constKeyProvider{key: []byte("legacy-boundary-key-material!!!"), version: 0}); err != nil { + t.Fatal(err) + } + + got, err := s.GetNode(ctx, "legacy-1") + if err != nil { + t.Fatal(err) + } + if got.Content != "written in plaintext" { + t.Fatalf("legacy row unreadable after enabling encryption: %q", got.Content) + } + + // Progressive re-encryption via update. + got.Content = "now encrypted" + if err := s.UpdateNode(ctx, got); err != nil { + t.Fatal(err) + } + var raw string + if err := s.DB().QueryRowContext(ctx, `SELECT content FROM nodes WHERE id=?`, "legacy-1").Scan(&raw); err != nil { + t.Fatal(err) + } + if !IsEncryptedValue(raw) { + t.Fatalf("legacy row not upgraded to ciphertext on update: %q", raw) + } +} + +// TestStoreEncryptionWrongKeyFailFast asserts the store refuses to serve +// ciphertext garbage when opened without the right key. +func TestStoreEncryptionWrongKeyFailFast(t *testing.T) { + s, cleanup := setupStore(t) + defer cleanup() + if err := s.EnableEncryption(&constKeyProvider{key: []byte("first-key-material-32-bytes!!!!"), version: 0}); err != nil { + t.Fatal(err) + } + ctx := context.Background() + if err := s.CreateNode(ctx, &Node{ + ID: "k1", Type: "convention", Content: "sealed with key one", + ContentHash: "h-k1", Scope: "project", Project: "test", + }); err != nil { + t.Fatal(err) + } + + // Reopen with a different key: decrypt must fail, not return garbage. + dbPath := s.dbPath + _ = s.Close() + + s2, err := NewStore(dbPath) + if err != nil { + t.Fatal(err) + } + defer s2.Close() + if err := s2.EnableEncryption(&constKeyProvider{key: []byte("second-key-material-32-bytes!!!"), version: 0}); err != nil { + t.Fatal(err) + } + if _, err := s2.GetNode(ctx, "k1"); !errors.Is(err, ErrDecryptionFailed) { + t.Fatalf("wrong key must fail reads, got %v", err) + } +} + +// TestStoreEncryptionInTx covers the transactional store path. +func TestStoreEncryptionInTx(t *testing.T) { + s, cleanup := setupStore(t) + defer cleanup() + if err := s.EnableEncryption(&constKeyProvider{key: []byte("tx-key-material-32-bytes!!!!!!!"), version: 0}); err != nil { + t.Fatal(err) + } + ctx := context.Background() + + err := s.WithTx(ctx, func(tx Storage) error { + if err := tx.CreateNode(ctx, &Node{ + ID: "tx-enc", Type: "convention", Content: "tx plaintext", + ContentHash: "h-tx", Scope: "project", Project: "test", + }); err != nil { + return err + } + got, err := tx.GetNode(ctx, "tx-enc") + if err != nil { + return err + } + if got.Content != "tx plaintext" { + return errors.New("tx read did not decrypt") + } + return nil + }) + if err != nil { + t.Fatal(err) + } + + var raw string + if err := s.DB().QueryRowContext(ctx, `SELECT content FROM nodes WHERE id=?`, "tx-enc").Scan(&raw); err != nil { + t.Fatal(err) + } + if !IsEncryptedValue(raw) { + t.Fatalf("tx write not encrypted at rest: %q", raw) + } +} + +// TestStoreNoEncryptionByDefault verifies default behaviour is unchanged: +// plaintext on disk, bookkeeping columns false/0. +func TestStoreNoEncryptionByDefault(t *testing.T) { + s, cleanup := setupStore(t) + defer cleanup() + ctx := context.Background() + if err := s.CreateNode(ctx, &Node{ + ID: "plain-1", Type: "convention", Content: "stays plaintext", + ContentHash: "h-plain", Scope: "project", Project: "test", + }); err != nil { + t.Fatal(err) + } + + var raw string + var encrypted bool + var keyVersion int + if err := s.DB().QueryRowContext(ctx, + `SELECT content, encrypted, encryption_key_version FROM nodes WHERE id=?`, "plain-1"). + Scan(&raw, &encrypted, &keyVersion); err != nil { + if errors.Is(err, sql.ErrNoRows) { + t.Fatal("node row missing") + } + t.Fatal(err) + } + if raw != "stays plaintext" || encrypted || keyVersion != 0 { + t.Fatalf("default store must stay plaintext: %q enc=%v kv=%d", raw, encrypted, keyVersion) + } +} + +// TestStoreEncryptionSearchFallback pins the FTS5 boundary: the FTS index +// stores ciphertext, so keyword search falls back to an in-memory scan over +// decrypted content when the cipher is active. +func TestStoreEncryptionSearchFallback(t *testing.T) { + s, cleanup := setupStore(t) + defer cleanup() + if err := s.EnableEncryption(&constKeyProvider{key: []byte("search-fallback-key-material!!!"), version: 0}); err != nil { + t.Fatal(err) + } + ctx := context.Background() + for i, content := range []string{ + "always run the linter before commit", + "deploy through the golden pipeline", + "keep functions under forty lines", + } { + if err := s.CreateNode(ctx, &Node{ + ID: fmt.Sprintf("search-%d", i), Type: "convention", Content: content, + ContentHash: fmt.Sprintf("h-search-%d", i), Scope: "project", Project: "test", + }); err != nil { + t.Fatal(err) + } + } + + nodes, err := s.SearchNodes(ctx, "linter", 5) + if err != nil { + t.Fatal(err) + } + if len(nodes) != 1 || nodes[0].Content != "always run the linter before commit" { + t.Fatalf("encrypted keyword search failed: %d hits, %+v", len(nodes), nodes) + } + + // Multi-token OR matching across content/summary. + nodes, err = s.SearchNodes(ctx, "pipeline forty", 5) + if err != nil { + t.Fatal(err) + } + if len(nodes) != 2 { + t.Fatalf("expected 2 OR-matched nodes, got %d", len(nodes)) + } + + // Limit respected. + nodes, err = s.SearchNodes(ctx, "the", 2) + if err != nil { + t.Fatal(err) + } + if len(nodes) != 2 { + t.Fatalf("limit 2 not honored, got %d", len(nodes)) + } +} + +// TestEncryptValuesNilCipher guards the hot-path passthrough. +func TestEncryptValuesNilCipher(t *testing.T) { + out, err := encryptValues(nil, "a", "b") + if err != nil || out[0] != "a" || out[1] != "b" { + t.Fatalf("nil cipher must pass through: %v %v", out, err) + } + out, err = decryptValues(nil, "x") + if err != nil || out[0] != "x" { + t.Fatalf("nil cipher decrypt passthrough: %v %v", out, err) + } +} diff --git a/storage/errors.go b/storage/errors.go index b108e6c..aad713d 100644 --- a/storage/errors.go +++ b/storage/errors.go @@ -23,4 +23,9 @@ var ( ErrBatchTooLarge = errors.New("batch size exceeds database limit") ErrVersionNotFound = errors.New("version not found") ErrEmbeddingNotFound = errors.New("embedding not found") + + ErrKeyMaterial = errors.New("encryption key material unavailable") + ErrUnknownKeyVersion = errors.New("unknown encryption key version") + ErrMalformedCiphertext = errors.New("malformed ciphertext") + ErrDecryptionFailed = errors.New("decryption failed") ) diff --git a/storage/node_crypto.go b/storage/node_crypto.go new file mode 100644 index 0000000..81c35a7 --- /dev/null +++ b/storage/node_crypto.go @@ -0,0 +1,94 @@ +package storage + +import "fmt" + +// NodeCipher methods below are written with a nil receiver guard so callers +// can pass a *NodeCipher that is nil when encryption is disabled and get +// plaintext passthrough without branching at every call site. + +// encryptValues seals the provided node content/summary with c. When c is nil +// (encryption disabled) or a value is empty, the value is returned unchanged. +// Encrypting an already-encrypted value is harmless but pointless; callers +// always pass fresh plaintext loaded from the API. +func encryptValues(c *NodeCipher, values ...string) ([]string, error) { + if c == nil { + out := make([]string, len(values)) + copy(out, values) + return out, nil + } + out := make([]string, len(values)) + for i, v := range values { + enc, err := c.Encrypt(v) + if err != nil { + return nil, fmt.Errorf("encrypt field %d: %w", i, err) + } + out[i] = enc + } + return out, nil +} + +// decryptValues opens the provided stored values with c. A nil c, or any value +// lacking the cipher prefix (legacy plaintext written before encryption was +// enabled), passes through unchanged. This is what makes a store transparently +// readable across the plaintext→encrypted boundary. +func decryptValues(c *NodeCipher, values ...string) ([]string, error) { + if c == nil { + out := make([]string, len(values)) + copy(out, values) + return out, nil + } + out := make([]string, len(values)) + for i, v := range values { + dec, err := c.Decrypt(v) + if err != nil { + return nil, fmt.Errorf("decrypt field %d: %w", i, err) + } + out[i] = dec + } + return out, nil +} + +// decryptNodeInplace seals nothing; it opens a freshly scanned Node's +// encrypted content/summary so the caller sees plaintext. On any field error +// the node is returned as-is and the error propagates — a tampered or +// key-rotated-away row must not silently degrade to ciphertext. +func decryptNodeInplace(n *Node, c *NodeCipher) error { + if n == nil || c == nil { + return nil + } + vals, err := decryptValues(c, n.Content, n.Summary) + if err != nil { + return err + } + n.Content = vals[0] + n.Summary = vals[1] + return nil +} + +// decryptVersionInplace opens a NodeVersion's stored content (version history +// is encrypted alongside the live content so history does not leak). +func decryptVersionInplace(v *NodeVersion, c *NodeCipher) error { + if v == nil || c == nil { + return nil + } + vals, err := decryptValues(c, v.Content) + if err != nil { + return err + } + v.Content = vals[0] + return nil +} + +// keyVersionForWrite reports the encryption key version stamped into +// nodes.encryption_key_version for new writes, or 0 when encryption is +// disabled. nodes.encrypted is set in lockstep. +func keyVersionForWrite(c *NodeCipher) int { + if c == nil { + return 0 + } + return c.KeyVersion() +} + +// encryptedFlag returns the nodes.encrypted value for a write: true when the +// cipher is active, false for plaintext. +func encryptedFlag(c *NodeCipher) bool { return c != nil } diff --git a/storage/prefix.go b/storage/prefix.go index a723960..e6cf68e 100644 --- a/storage/prefix.go +++ b/storage/prefix.go @@ -22,5 +22,5 @@ func (s *Store) FindByPrefix(ctx context.Context, prefix string) ([]*Node, error return nil, err } defer func() { _ = rows.Close() }() - return scanNodes(rows) + return scanNodes(rows, s.cipher()) } diff --git a/storage/sqlite.go b/storage/sqlite.go index aae42b6..6dd6bdd 100644 --- a/storage/sqlite.go +++ b/storage/sqlite.go @@ -230,11 +230,39 @@ const defaultBusyTimeoutMs = 30000 type Store struct { db *sql.DB + dbPath string cache *stmtCache queryTimeout time.Duration processLock *ProcessLock hnswMu sync.Mutex hnswIndexes map[string]*HNSWIndex // keyed by embedding model + + cipherMu sync.RWMutex + nodeCipher *NodeCipher // nil until EnableEncryption; guards progressive activation +} + +// EnableEncryption activates at-rest encryption of node content/summary using +// keys from the given provider. Call it right after NewStore, before any +// reads, so no plaintext can be served by the API once writes start +// encrypting. It is idempotent: the first call wins; later calls are no-ops. +func (s *Store) EnableEncryption(p KeyProvider) error { + c, err := NewNodeCipher(p) + if err != nil { + return err + } + s.cipherMu.Lock() + defer s.cipherMu.Unlock() + if s.nodeCipher == nil { + s.nodeCipher = c + } + return nil +} + +// cipher returns the active node cipher, or nil when encryption is disabled. +func (s *Store) cipher() *NodeCipher { + s.cipherMu.RLock() + defer s.cipherMu.RUnlock() + return s.nodeCipher } // DB returns the underlying database connection for direct queries. @@ -284,7 +312,7 @@ func NewStore(dbPath string) (*Store, error) { if err := applyPragmas(context.Background(), db); err != nil { return nil, err } - s := &Store{db: db, queryTimeout: 30 * time.Second, hnswIndexes: make(map[string]*HNSWIndex)} + s := &Store{db: db, dbPath: dbPath, queryTimeout: 30 * time.Second, hnswIndexes: make(map[string]*HNSWIndex)} if err := s.createTables(); err != nil { return nil, err } @@ -372,11 +400,11 @@ func (s *Store) SearchNodeByHash(ctx context.Context, hash, scope, project strin return retryOnBusyVal(func() (*Node, error) { ctx, cancel := s.withTimeout(ctx) defer cancel() - return searchNodeByHashQ(ctx, s.q(), hash, scope, project) + return searchNodeByHashQ(ctx, s.q(), hash, scope, project, s.cipher()) }, 5, 50*time.Millisecond) } -func searchNodeByHashQ(ctx context.Context, q queryable, hash, scope, project string) (*Node, error) { +func searchNodeByHashQ(ctx context.Context, q queryable, hash, scope, project string, c *NodeCipher) (*Node, error) { n := &Node{} var at sql.NullTime var key sql.NullString @@ -399,6 +427,9 @@ func searchNodeByHashQ(ctx context.Context, q queryable, hash, scope, project st if key.Valid { n.Key = key.String } + if err := decryptNodeInplace(n, c); err != nil { + return nil, err + } return n, nil } @@ -658,7 +689,7 @@ func (s *Store) GetNodesByFile(ctx context.Context, filePath string) ([]*Node, e return nil, err } defer func() { _ = rows.Close() }() - return scanNodes(rows) + return scanNodes(rows, s.cipher()) }, 5, 50*time.Millisecond) } @@ -678,7 +709,7 @@ func nullString(s string) sql.NullString { return sql.NullString{String: s, Valid: true} } -func scanNodes(rows *sql.Rows) ([]*Node, error) { +func scanNodes(rows *sql.Rows, c *NodeCipher) ([]*Node, error) { var out []*Node for rows.Next() { n := &Node{} @@ -693,6 +724,9 @@ func scanNodes(rows *sql.Rows) ([]*Node, error) { if key.Valid { n.Key = key.String } + if err := decryptNodeInplace(n, c); err != nil { + return nil, err + } out = append(out, n) } return out, rows.Err() diff --git a/storage/sqlite_edges.go b/storage/sqlite_edges.go index 3859ef6..abe2e7f 100644 --- a/storage/sqlite_edges.go +++ b/storage/sqlite_edges.go @@ -242,18 +242,18 @@ func (s *Store) GetNeighbors(ctx context.Context, nodeID string) ([]*Node, error return retryOnBusyVal(func() ([]*Node, error) { ctx, cancel := s.withTimeout(ctx) defer cancel() - return getNeighborsQ(ctx, s.q(), nodeID) + return getNeighborsQ(ctx, s.q(), nodeID, s.cipher()) }, 5, 50*time.Millisecond) } -func getNeighborsQ(ctx context.Context, q queryable, nodeID string) ([]*Node, error) { +func getNeighborsQ(ctx context.Context, q queryable, nodeID string, c *NodeCipher) ([]*Node, error) { rows, err := q.QueryContext(ctx, `SELECT DISTINCT n.id, n.type, n.content, n.content_hash, n.summary, n.scope, n.project, n.tier, n.tags, n.key, n.pinned, n.confidence, n.access_count, n.created_at, n.updated_at, n.accessed_at, n.source_session, n.source_agent, n.version FROM nodes n JOIN edges e ON (e.to_id = n.id AND e.from_id = ?) OR (e.from_id = n.id AND e.to_id = ?)`, nodeID, nodeID) if err != nil { return nil, err } defer func() { _ = rows.Close() }() - return scanNodes(rows) + return scanNodes(rows, c) } // CountEdges returns inbound and outbound edge counts for a node. diff --git a/storage/sqlite_nodes.go b/storage/sqlite_nodes.go index 90a8b1f..7dba1b4 100644 --- a/storage/sqlite_nodes.go +++ b/storage/sqlite_nodes.go @@ -9,6 +9,7 @@ import ( "database/sql" "errors" "fmt" + "sort" "strings" "time" @@ -24,7 +25,7 @@ func (s *Store) CreateNode(ctx context.Context, n *Node) error { ctx, cancel := s.withTimeout(ctx) defer cancel() err := retryOnBusy(func() error { - return createNodeQ(ctx, s.q(), n) + return createNodeQ(ctx, s.q(), n, s.cipher()) }, 5, 2*time.Millisecond) attrs := attribute.NewSet(attribute.String("op", "create_node")) telemetry.SQLiteQueryDuration.Record(ctx, time.Since(start).Seconds(), metric.WithAttributeSet(attrs)) @@ -32,11 +33,15 @@ func (s *Store) CreateNode(ctx context.Context, n *Node) error { return err } -func createNodeQ(ctx context.Context, q queryable, n *Node) error { - _, err := q.ExecContext(ctx, `INSERT INTO nodes (id, type, content, content_hash, summary, scope, project, tier, tags, key, pinned, confidence, access_count, created_at, updated_at, accessed_at, source_session, source_agent, version) - VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, - n.ID, n.Type, n.Content, n.ContentHash, n.Summary, n.Scope, n.Project, n.Tier, n.Tags, nullString(n.Key), n.Pinned, n.Confidence, n.AccessCount, - n.CreatedAt, n.UpdatedAt, nullTime(n.AccessedAt), n.SourceSession, n.SourceAgent, n.Version) +func createNodeQ(ctx context.Context, q queryable, n *Node, c *NodeCipher) error { + enc, err := encryptValues(c, n.Content, n.Summary) + if err != nil { + return err + } + _, err = q.ExecContext(ctx, `INSERT INTO nodes (id, type, content, content_hash, summary, scope, project, tier, tags, key, pinned, confidence, access_count, created_at, updated_at, accessed_at, source_session, source_agent, version, encrypted, encryption_key_version) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + n.ID, n.Type, enc[0], n.ContentHash, enc[1], n.Scope, n.Project, n.Tier, n.Tags, nullString(n.Key), n.Pinned, n.Confidence, n.AccessCount, + n.CreatedAt, n.UpdatedAt, nullTime(n.AccessedAt), n.SourceSession, n.SourceAgent, n.Version, encryptedFlag(c), keyVersionForWrite(c)) if err != nil && strings.Contains(err.Error(), "UNIQUE constraint failed") { return fmt.Errorf("%w: %s", ErrDuplicateNode, err) } @@ -73,7 +78,7 @@ func (s *Store) GetNode(ctx context.Context, id string) (*Node, error) { n, err := retryOnBusyVal(func() (*Node, error) { ctx, cancel := s.withTimeout(ctx) defer cancel() - return getNodeQ(ctx, s.q(), id) + return getNodeQ(ctx, s.q(), id, s.cipher()) }, 5, 50*time.Millisecond) attrs := attribute.NewSet(attribute.String("op", "get_node")) telemetry.SQLiteQueryDuration.Record(ctx, time.Since(start).Seconds(), metric.WithAttributeSet(attrs)) @@ -82,7 +87,7 @@ func (s *Store) GetNode(ctx context.Context, id string) (*Node, error) { return n, err } -func getNodeQ(ctx context.Context, q queryable, id string) (*Node, error) { +func getNodeQ(ctx context.Context, q queryable, id string, c *NodeCipher) (*Node, error) { n := &Node{} var accessedAt sql.NullTime var key sql.NullString @@ -100,6 +105,9 @@ func getNodeQ(ctx context.Context, q queryable, id string) (*Node, error) { if key.Valid { n.Key = key.String } + if err := decryptNodeInplace(n, c); err != nil { + return nil, err + } return n, nil } @@ -108,10 +116,10 @@ func getNodeQ(ctx context.Context, q queryable, id string) (*Node, error) { func (s *Store) GetNodeByKey(ctx context.Context, key, project string) (*Node, error) { ctx, cancel := s.withTimeout(ctx) defer cancel() - return getNodeByKeyQ(ctx, s.q(), key, project) + return getNodeByKeyQ(ctx, s.q(), key, project, s.cipher()) } -func getNodeByKeyQ(ctx context.Context, q queryable, key, project string) (*Node, error) { +func getNodeByKeyQ(ctx context.Context, q queryable, key, project string, c *NodeCipher) (*Node, error) { n := &Node{} var accessedAt sql.NullTime var k sql.NullString @@ -129,6 +137,9 @@ func getNodeByKeyQ(ctx context.Context, q queryable, key, project string) (*Node if k.Valid { n.Key = k.String } + if err := decryptNodeInplace(n, c); err != nil { + return nil, err + } return n, nil } @@ -137,7 +148,7 @@ func (s *Store) UpdateNode(ctx context.Context, n *Node) error { err := retryOnBusy(func() error { ctx, cancel := s.withTimeout(ctx) defer cancel() - return updateNodeQ(ctx, s.q(), n) + return updateNodeQ(ctx, s.q(), n, s.cipher()) }, 5, 50*time.Millisecond) attrs := attribute.NewSet(attribute.String("op", "update_node")) telemetry.SQLiteQueryDuration.Record(ctx, time.Since(start).Seconds(), metric.WithAttributeSet(attrs)) @@ -145,10 +156,14 @@ func (s *Store) UpdateNode(ctx context.Context, n *Node) error { return err } -func updateNodeQ(ctx context.Context, q queryable, n *Node) error { - _, err := q.ExecContext(ctx, `UPDATE nodes SET type=?, content=?, content_hash=?, summary=?, scope=?, project=?, tier=?, tags=?, key=?, pinned=?, confidence=?, access_count=?, updated_at=?, accessed_at=?, source_session=?, source_agent=?, version=? WHERE id=?`, - n.Type, n.Content, n.ContentHash, n.Summary, n.Scope, n.Project, n.Tier, n.Tags, nullString(n.Key), n.Pinned, n.Confidence, n.AccessCount, - n.UpdatedAt, nullTime(n.AccessedAt), n.SourceSession, n.SourceAgent, n.Version, n.ID) +func updateNodeQ(ctx context.Context, q queryable, n *Node, c *NodeCipher) error { + enc, err := encryptValues(c, n.Content, n.Summary) + if err != nil { + return err + } + _, err = q.ExecContext(ctx, `UPDATE nodes SET type=?, content=?, content_hash=?, summary=?, scope=?, project=?, tier=?, tags=?, key=?, pinned=?, confidence=?, access_count=?, updated_at=?, accessed_at=?, source_session=?, source_agent=?, version=?, encrypted=?, encryption_key_version=? WHERE id=?`, + n.Type, enc[0], n.ContentHash, enc[1], n.Scope, n.Project, n.Tier, n.Tags, nullString(n.Key), n.Pinned, n.Confidence, n.AccessCount, + n.UpdatedAt, nullTime(n.AccessedAt), n.SourceSession, n.SourceAgent, n.Version, encryptedFlag(c), keyVersionForWrite(c), n.ID) if err != nil { return err } @@ -177,12 +192,16 @@ func (s *Store) UpdateNodeContent(ctx context.Context, id, newContent string) er return retryOnBusy(func() error { ctx, cancel := s.withTimeout(ctx) defer cancel() - return updateNodeContentQ(ctx, s.q(), id, newContent) + return updateNodeContentQ(ctx, s.q(), id, newContent, s.cipher()) }, 5, 50*time.Millisecond) } -func updateNodeContentQ(ctx context.Context, q queryable, id, newContent string) error { - _, err := q.ExecContext(ctx, `UPDATE nodes SET content=?, updated_at=CURRENT_TIMESTAMP WHERE id=?`, newContent, id) +func updateNodeContentQ(ctx context.Context, q queryable, id, newContent string, c *NodeCipher) error { + enc, err := encryptValues(c, newContent) + if err != nil { + return err + } + _, err = q.ExecContext(ctx, `UPDATE nodes SET content=?, updated_at=CURRENT_TIMESTAMP, encrypted=?, encryption_key_version=? WHERE id=?`, enc[0], encryptedFlag(c), keyVersionForWrite(c), id) return err } @@ -215,7 +234,7 @@ func (s *Store) ListNodes(ctx context.Context, f NodeFilter) ([]*Node, error) { nodes, err := retryOnBusyVal(func() ([]*Node, error) { ctx, cancel := s.withTimeout(ctx) defer cancel() - return listNodesQ(ctx, s.q(), f) + return listNodesQ(ctx, s.q(), f, s.cipher()) }, 5, 50*time.Millisecond) attrs := attribute.NewSet(attribute.String("op", "list_nodes")) telemetry.SQLiteQueryDuration.Record(ctx, time.Since(start).Seconds(), metric.WithAttributeSet(attrs)) @@ -223,7 +242,7 @@ func (s *Store) ListNodes(ctx context.Context, f NodeFilter) ([]*Node, error) { return nodes, err } -func listNodesQ(ctx context.Context, q queryable, f NodeFilter) ([]*Node, error) { +func listNodesQ(ctx context.Context, q queryable, f NodeFilter, c *NodeCipher) ([]*Node, error) { query := "SELECT id, type, content, content_hash, summary, scope, project, tier, tags, key, pinned, confidence, access_count, created_at, updated_at, accessed_at, source_session, source_agent, version FROM nodes WHERE 1=1" var args []any if f.Type != "" { @@ -286,7 +305,7 @@ func listNodesQ(ctx context.Context, q queryable, f NodeFilter) ([]*Node, error) return nil, err } defer func() { _ = rows.Close() }() - return scanNodes(rows) + return scanNodes(rows, c) } // escapeFTS5 escapes special FTS5 characters by wrapping each token in double @@ -307,7 +326,7 @@ func (s *Store) SearchNodes(ctx context.Context, query string, limit int) ([]*No nodes, err := retryOnBusyVal(func() ([]*Node, error) { ctx, cancel := s.withTimeout(ctx) defer cancel() - return searchNodesQ(ctx, s.q(), query, limit) + return searchNodesQ(ctx, s.q(), query, limit, s.cipher()) }, 5, 50*time.Millisecond) attrs := attribute.NewSet(attribute.String("op", "search_nodes")) telemetry.SQLiteQueryDuration.Record(ctx, time.Since(start).Seconds(), metric.WithAttributeSet(attrs)) @@ -315,10 +334,17 @@ func (s *Store) SearchNodes(ctx context.Context, query string, limit int) ([]*No return nodes, err } -func searchNodesQ(ctx context.Context, q queryable, query string, limit int) ([]*Node, error) { +func searchNodesQ(ctx context.Context, q queryable, query string, limit int, c *NodeCipher) ([]*Node, error) { if limit <= 0 { limit = 10 } + // FTS5 indexes the stored nodes.content column — ciphertext when + // encryption is active — so keyword search would match nothing. Fall + // back to an in-memory token scan over decrypted content; fine at + // yaad's graph sizes and keeps Recall working with encryption on. + if c != nil { + return searchNodesInMemory(ctx, q, query, limit, c) + } ftsQuery := escapeFTS5(query) rows, err := q.QueryContext(ctx, `SELECT n.id, n.type, n.content, n.content_hash, n.summary, n.scope, n.project, n.tier, n.tags, n.key, n.pinned, n.confidence, n.access_count, n.created_at, n.updated_at, n.accessed_at, n.source_session, n.source_agent, n.version FROM nodes_fts f JOIN nodes n ON f.rowid = n.rowid WHERE nodes_fts MATCH ? ORDER BY rank LIMIT ?`, ftsQuery, limit) @@ -326,7 +352,63 @@ func searchNodesQ(ctx context.Context, q queryable, query string, limit int) ([] return nil, err } defer func() { _ = rows.Close() }() - return scanNodes(rows) + return scanNodes(rows, c) +} + +// searchNodesInMemory scans all nodes, decrypting as it goes, and matches the +// whitespace-separated query tokens (case-insensitive substring, OR semantics +// like escapeFTS5) against content and summary. Results are ordered by the +// number of distinct tokens matched (descending), then most-recently-updated. +func searchNodesInMemory(ctx context.Context, q queryable, query string, limit int, c *NodeCipher) ([]*Node, error) { + tokens := strings.Fields(query) + if len(tokens) == 0 { + return nil, nil + } + for i := range tokens { + tokens[i] = strings.ToLower(tokens[i]) + } + + rows, err := q.QueryContext(ctx, `SELECT id, type, content, content_hash, summary, scope, project, tier, tags, key, pinned, confidence, access_count, created_at, updated_at, accessed_at, source_session, source_agent, version FROM nodes`) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + all, err := scanNodes(rows, c) + if err != nil { + return nil, err + } + + type scored struct { + node *Node + score int + } + var hits []scored + for _, n := range all { + hay := strings.ToLower(n.Content + "\n" + n.Summary) + matched := 0 + for _, tok := range tokens { + if strings.Contains(hay, tok) { + matched++ + } + } + if matched > 0 { + hits = append(hits, scored{n, matched}) + } + } + sort.Slice(hits, func(i, j int) bool { + if hits[i].score != hits[j].score { + return hits[i].score > hits[j].score + } + return hits[i].node.UpdatedAt.After(hits[j].node.UpdatedAt) + }) + if len(hits) > limit { + hits = hits[:limit] + } + out := make([]*Node, 0, len(hits)) + for _, h := range hits { + out = append(out, h.node) + } + return out, nil } // --- Versions --- @@ -340,21 +422,25 @@ func (s *Store) SaveVersion(ctx context.Context, nodeID string, content, changed return err } defer func() { _ = tx.Rollback() }() - if err := saveVersionQ(ctx, tx, nodeID, content, changedBy, reason); err != nil { + if err := saveVersionQ(ctx, tx, nodeID, content, changedBy, reason, s.cipher()); err != nil { return err } return tx.Commit() }, 5, 50*time.Millisecond) } -func saveVersionQ(ctx context.Context, q queryable, nodeID string, content, changedBy, reason string) error { +func saveVersionQ(ctx context.Context, q queryable, nodeID string, content, changedBy, reason string, c *NodeCipher) error { + enc, err := encryptValues(c, content) + if err != nil { + return err + } var maxVer int - err := q.QueryRowContext(ctx, `SELECT COALESCE(MAX(version), 0) FROM node_versions WHERE node_id=?`, nodeID).Scan(&maxVer) + err = q.QueryRowContext(ctx, `SELECT COALESCE(MAX(version), 0) FROM node_versions WHERE node_id=?`, nodeID).Scan(&maxVer) if err != nil { return err } _, err = q.ExecContext(ctx, `INSERT INTO node_versions (node_id, version, content, changed_at, changed_by, reason) VALUES (?, ?, ?, ?, ?, ?)`, - nodeID, maxVer+1, content, time.Now().UTC(), changedBy, reason) + nodeID, maxVer+1, enc[0], time.Now().UTC(), changedBy, reason) return err } @@ -362,11 +448,11 @@ func (s *Store) GetVersions(ctx context.Context, nodeID string) ([]*NodeVersion, return retryOnBusyVal(func() ([]*NodeVersion, error) { ctx, cancel := s.withTimeout(ctx) defer cancel() - return getVersionsQ(ctx, s.q(), nodeID) + return getVersionsQ(ctx, s.q(), nodeID, s.cipher()) }, 5, 50*time.Millisecond) } -func getVersionsQ(ctx context.Context, q queryable, nodeID string) ([]*NodeVersion, error) { +func getVersionsQ(ctx context.Context, q queryable, nodeID string, c *NodeCipher) ([]*NodeVersion, error) { rows, err := q.QueryContext(ctx, `SELECT node_id, version, content, changed_at, changed_by, reason FROM node_versions WHERE node_id=? ORDER BY version`, nodeID) if err != nil { return nil, err @@ -378,6 +464,9 @@ func getVersionsQ(ctx context.Context, q queryable, nodeID string) ([]*NodeVersi if err := rows.Scan(&v.NodeID, &v.Version, &v.Content, &v.ChangedAt, &v.ChangedBy, &v.Reason); err != nil { return nil, err } + if err := decryptVersionInplace(v, c); err != nil { + return nil, err + } out = append(out, v) } return out, rows.Err() @@ -592,11 +681,11 @@ func (s *Store) GetNodesBatch(ctx context.Context, ids []string) ([]*Node, error return retryOnBusyVal(func() ([]*Node, error) { ctx, cancel := s.withTimeout(ctx) defer cancel() - return getNodesBatchQ(ctx, s.q(), ids) + return getNodesBatchQ(ctx, s.q(), ids, s.cipher()) }, 5, 50*time.Millisecond) } -func getNodesBatchQ(ctx context.Context, q queryable, ids []string) ([]*Node, error) { +func getNodesBatchQ(ctx context.Context, q queryable, ids []string, c *NodeCipher) ([]*Node, error) { if len(ids) == 0 { return nil, nil } @@ -618,7 +707,7 @@ func getNodesBatchQ(ctx context.Context, q queryable, ids []string) ([]*Node, er if err != nil { return nil, err } - nodes, err := scanNodes(rows) + nodes, err := scanNodes(rows, c) _ = rows.Close() if err != nil { return nil, err diff --git a/storage/sqlite_tx.go b/storage/sqlite_tx.go index 0a8bcd9..22a5c63 100644 --- a/storage/sqlite_tx.go +++ b/storage/sqlite_tx.go @@ -53,24 +53,28 @@ type txStore struct { } // txStore is a thin wrapper that delegates all operations to shared *Q functions. -func (t *txStore) CreateNode(ctx context.Context, n *Node) error { return createNodeQ(ctx, t.tx, n) } +func (t *txStore) CreateNode(ctx context.Context, n *Node) error { + return createNodeQ(ctx, t.tx, n, t.store.cipher()) +} func (t *txStore) GetNode(ctx context.Context, id string) (*Node, error) { - return getNodeQ(ctx, t.tx, id) + return getNodeQ(ctx, t.tx, id, t.store.cipher()) } func (t *txStore) GetNodeByKey(ctx context.Context, key, project string) (*Node, error) { - return getNodeByKeyQ(ctx, t.tx, key, project) + return getNodeByKeyQ(ctx, t.tx, key, project, t.store.cipher()) } func (t *txStore) GetNodesBatch(ctx context.Context, ids []string) ([]*Node, error) { - return getNodesBatchQ(ctx, t.tx, ids) + return getNodesBatchQ(ctx, t.tx, ids, t.store.cipher()) } -func (t *txStore) UpdateNode(ctx context.Context, n *Node) error { return updateNodeQ(ctx, t.tx, n) } +func (t *txStore) UpdateNode(ctx context.Context, n *Node) error { + return updateNodeQ(ctx, t.tx, n, t.store.cipher()) +} func (t *txStore) UpdateNodeContent(ctx context.Context, id, newContent string) error { - return updateNodeContentQ(ctx, t.tx, id, newContent) + return updateNodeContentQ(ctx, t.tx, id, newContent, t.store.cipher()) } func (t *txStore) ArchiveNode(ctx context.Context, id string) (bool, error) { @@ -80,19 +84,19 @@ func (t *txStore) ArchiveNode(ctx context.Context, id string) (bool, error) { func (t *txStore) DeleteNode(ctx context.Context, id string) error { return deleteNodeQ(ctx, t.tx, id) } func (t *txStore) ListNodes(ctx context.Context, f NodeFilter) ([]*Node, error) { - return listNodesQ(ctx, t.tx, f) + return listNodesQ(ctx, t.tx, f, t.store.cipher()) } func (t *txStore) SearchNodes(ctx context.Context, query string, limit int) ([]*Node, error) { - return searchNodesQ(ctx, t.tx, query, limit) + return searchNodesQ(ctx, t.tx, query, limit, t.store.cipher()) } func (t *txStore) SearchNodeByHash(ctx context.Context, hash, scope, project string) (*Node, error) { - return searchNodeByHashQ(ctx, t.tx, hash, scope, project) + return searchNodeByHashQ(ctx, t.tx, hash, scope, project, t.store.cipher()) } func (t *txStore) GetNeighbors(ctx context.Context, nodeID string) ([]*Node, error) { - return getNeighborsQ(ctx, t.tx, nodeID) + return getNeighborsQ(ctx, t.tx, nodeID, t.store.cipher()) } func (t *txStore) CreateEdge(ctx context.Context, e *Edge) error { return createEdgeQ(ctx, t.tx, e) } @@ -176,11 +180,11 @@ func (t *txStore) ListSessions(ctx context.Context, project string, limit int) ( } func (t *txStore) SaveVersion(ctx context.Context, nodeID string, content, changedBy, reason string) error { - return saveVersionQ(ctx, t.tx, nodeID, content, changedBy, reason) + return saveVersionQ(ctx, t.tx, nodeID, content, changedBy, reason, t.store.cipher()) } func (t *txStore) GetVersions(ctx context.Context, nodeID string) ([]*NodeVersion, error) { - return getVersionsQ(ctx, t.tx, nodeID) + return getVersionsQ(ctx, t.tx, nodeID, t.store.cipher()) } func (t *txStore) SaveEmbedding(ctx context.Context, nodeID, model string, vector []float32) error { @@ -291,6 +295,13 @@ func (s *Store) RollbackToVersion(ctx context.Context, nodeID string, version in } else { return fmt.Errorf("unexpected storage type in transaction") } + // Version history stores ciphertext; decrypt before re-encrypting + // via UpdateNodeContent / SaveVersion to avoid double-encryption. + vals, err := decryptValues(s.cipher(), content) + if err != nil { + return err + } + content = vals[0] if err := tx.UpdateNodeContent(ctx, nodeID, content); err != nil { return err } @@ -312,5 +323,10 @@ func (s *Store) DiffVersions(ctx context.Context, nodeID string, v1, v2 int) (co if err != nil { return "", "", fmt.Errorf("version %d not found: %w", v2, err) } - return content1, content2, nil + // Version history stores ciphertext; decrypt for a plaintext comparison. + vals, err := decryptValues(s.cipher(), content1, content2) + if err != nil { + return "", "", err + } + return vals[0], vals[1], nil } From 5b2d541743fc07e6e761c34ab80446c79ff143c4 Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Sat, 15 Aug 2026 20:48:26 +0530 Subject: [PATCH 2/3] docs: mark Phase 2 complete (PRs #53-59 incl. encryption) --- plans/PHASE1-SQLITE-HARDENING.md | 9 +++++---- 1 file changed, 5 insertions(+), 4 deletions(-) diff --git a/plans/PHASE1-SQLITE-HARDENING.md b/plans/PHASE1-SQLITE-HARDENING.md index 08ef052..9d29106 100644 --- a/plans/PHASE1-SQLITE-HARDENING.md +++ b/plans/PHASE1-SQLITE-HARDENING.md @@ -1,6 +1,6 @@ # Feature Specification: Phase 1 SQLite Hardening -**Status:** Implemented — Phase 1 merged via PRs #47–#51; Phase 2 items 1–3 merged via PRs #53–#55 +**Status:** Implemented — Phase 1 merged via PRs #47–#51; Phase 2 complete (PRs #53–#59: HNSW incremental/restore, backup rotation/scheduler, lock fix, encryption) **Author:** Patel230 **Date:** 2026-08-15 **Repos affected:** `GrayCodeAI/yaad` (via hawk submodule `external/yaad`) @@ -185,11 +185,12 @@ does not exist on Windows):** - [x] `feat/storage-hnsw` — PR #50 (`6c03254`) - [x] `feat/storage-backup` — PR #51 (`dfedbce`) -### Phase 2 (items 1-3 shipped) +### Phase 2 (items 1-4 shipped) - ~~Incremental HNSW updates on `SaveEmbedding`/`DeleteEmbedding`~~ — **done** (PR #53: `HNSWIndex.Upsert`/`Remove` with link pruning + entry-point repair; `Store` write paths apply them; transactional writes invalidate the cache after commit) - ~~Persist-and-restore of the HNSW graph across restarts~~ — **done** (PR #54: versioned payload in `embeddings_hnsw` carrying `version`/`m`/`efConstruction`; `HNSWIndex.Restore` requires parameter + membership parity, `BuildHNSWIndex` is the force-rebuild path) -- ~~Backup rotation~~ — **done** (PR #55: `Store.RotateBackups(dir, keep, maxAge)` with count/age retention and stale-`.tmp` cleanup; a daemon-side scheduler is not yet implemented) -- Not yet implemented: encryption columns (migration v5) wired to a real key provider +- ~~Backup rotation~~ — **done** (PR #55: `Store.RotateBackups(dir, keep, maxAge)` with count/age retention and stale-`.tmp` cleanup) +- ~~Backup scheduler~~ — **done** (PR #57: `Store.ScheduleBackups(dir, interval, keep, maxAge)` returns a `BackupScheduler` — immediate first snapshot, then per-interval best-effort `VACUUM INTO` + rotation; hawk starts it from session startup via `YaadBridge.EnsureBackups()`, hourly, keep 7 / 30d) +- ~~Encryption columns (migration v5) wired to a real key provider~~ — **done** (PR #59: app-layer AES-256-GCM — `KeyProvider` interface + `EnvKeyProvider` (`YAAD_ENCRYPTION_KEY`), version-bound ciphertexts `yaad.aes256gcm.v{n}`, transparent legacy-plaintext reads with progressive re-encrypt on update, `SearchNodes` in-memory fallback over decrypted content; hawk opts in via the same env var) ## Testing Strategy From f221983ebc25089e9648bc8a72537546c6b43572 Mon Sep 17 00:00:00 2001 From: Lakshman Patel Date: Sat, 15 Aug 2026 21:01:30 +0530 Subject: [PATCH 3/3] fix(tui): sync nested module deps after core x/crypto bump --- cmd/yaad-tui/go.mod | 5 +++-- cmd/yaad-tui/go.sum | 14 ++++++++------ 2 files changed, 11 insertions(+), 8 deletions(-) diff --git a/cmd/yaad-tui/go.mod b/cmd/yaad-tui/go.mod index 82c3931..ceb0cb0 100644 --- a/cmd/yaad-tui/go.mod +++ b/cmd/yaad-tui/go.mod @@ -1,6 +1,6 @@ module github.com/GrayCodeAI/yaad/cmd/yaad-tui -go 1.26.5 +go 1.26.6 require ( github.com/GrayCodeAI/yaad v0.0.0 @@ -48,8 +48,9 @@ require ( go.opentelemetry.io/otel v1.44.0 // indirect go.opentelemetry.io/otel/metric v1.44.0 // indirect go.opentelemetry.io/otel/trace v1.44.0 // indirect + golang.org/x/crypto v0.55.0 // indirect golang.org/x/sys v0.47.0 // indirect - golang.org/x/text v0.40.0 // indirect + golang.org/x/text v0.41.0 // indirect google.golang.org/protobuf v1.36.11 // indirect modernc.org/libc v1.72.5 // indirect modernc.org/mathutil v1.7.1 // indirect diff --git a/cmd/yaad-tui/go.sum b/cmd/yaad-tui/go.sum index 50a6cf8..c125995 100644 --- a/cmd/yaad-tui/go.sum +++ b/cmd/yaad-tui/go.sum @@ -101,19 +101,21 @@ go.uber.org/goleak v1.3.0 h1:2K3zAYmnTNqV73imy9J1T3WC+gmCePx2hEGkimedGto= go.uber.org/goleak v1.3.0/go.mod h1:CoHD4mav9JJNrW/WLlf7HGZPjdw8EucARQHekz1X6bE= go.yaml.in/yaml/v2 v2.4.4 h1:tuyd0P+2Ont/d6e2rl3be67goVK4R6deVxCUX5vyPaQ= go.yaml.in/yaml/v2 v2.4.4/go.mod h1:gMZqIpDtDqOfM0uNfy0SkpRhvUryYH0Z6wdMYcacYXQ= +golang.org/x/crypto v0.55.0 h1:+KWHjbgOaAQ66dh/YlkZKHlz9ZUlq61AFirAR9ntP8M= +golang.org/x/crypto v0.55.0/go.mod h1:uq0V9dE/fzQuJtbnL+2EhWOE63vo164FY8xqEnV9xis= golang.org/x/exp v0.0.0-20231006140011-7918f672742d h1:jtJma62tbqLibJ5sFQz8bKtEM8rJBtfilJ2qTU199MI= golang.org/x/exp v0.0.0-20231006140011-7918f672742d/go.mod h1:ldy0pHrwJyGW56pPQzzkH36rKxoZW1tw7ZJpeKx+hdo= -golang.org/x/mod v0.37.0 h1:vF1DjpVEshcIqoEaauuHebaLk1O1forxjxBaVn884JQ= -golang.org/x/mod v0.37.0/go.mod h1:m8S8VeM9r4dzDwjrKO0a1sZP3YjeMamRRlD+fmR2Q/0= +golang.org/x/mod v0.38.0 h1:MECBjubtXD7yj4HrhIUcywNaGeNVUdfVnxmPajOk4yk= +golang.org/x/mod v0.38.0/go.mod h1:V6Xz0pq8TQ3dGqVQ1FVHuelZpAL0uNhSkk9ogYP3c40= golang.org/x/sync v0.22.0 h1:SZjpbeLmrCk4xhRSZFNZW5gFUeCeFgjekvI/+gfScek= golang.org/x/sync v0.22.0/go.mod h1:9xrNwdLfx4jkKbNva9FpL6vEN7evnE43NNNJQ2LF3+0= golang.org/x/sys v0.0.0-20210809222454-d867a43fc93e/go.mod h1:oPkhp1MJrh7nUepCBck5+mAzfO9JrbApNNgaTdGDITg= golang.org/x/sys v0.47.0 h1:o7XGOvZQCADBQQ4Y7VNq2dRWQR7JmOUW8Kxx4ZsNgWs= golang.org/x/sys v0.47.0/go.mod h1:4GL1E5IUh+htKOUEOaiffhrAeqysfVGipDYzABqnCmw= -golang.org/x/text v0.40.0 h1:Ub2Z6/xjgF1WrYQz2nuITOEegKFtiIy+rieRJ5lHZKs= -golang.org/x/text v0.40.0/go.mod h1:hpnzDAfGV753zIKo+wk3u1bVKCGPbrnF7+7LBF/UHVY= -golang.org/x/tools v0.47.0 h1:7Kn5x/d1svx/PzryTsqeoZN4TZwqeH5pGWjefhLi/1Q= -golang.org/x/tools v0.47.0/go.mod h1:dFHnyTvFWY212G+h7ZY4Vsp/K3U4/7W9TyVaAul8uCA= +golang.org/x/text v0.41.0 h1:vz/seA0lnX87Othu2f/0L24RcgrXD9/YFTSuGjj3rH8= +golang.org/x/text v0.41.0/go.mod h1:jvf1O8ajNzZqhSrQBPbutR/EB83Cc0CFrezNQIwbb5M= +golang.org/x/tools v0.48.0 h1:3+hClM1aLL5mjMKm5ovokw9epgRXPuu2tILgismM6RE= +golang.org/x/tools v0.48.0/go.mod h1:08xX0orndb/F7jJxGDicx061tyd5pcMto75YMAXr6lk= google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/yaml.v3 v3.0.1 h1:fxVm/GzAzEWqLHuvctI91KS9hhNmmWOoWu0XTYJS7CA=