diff --git a/internal/cli/auth.go b/internal/cli/auth.go index f7064e917..d744e34bd 100644 --- a/internal/cli/auth.go +++ b/internal/cli/auth.go @@ -5,6 +5,7 @@ import ( "fmt" "io" "net/http" + "os" "strings" "time" @@ -124,9 +125,24 @@ func validateAuthFlags(sub string, a authArgs) error { // newAuthManager builds an oauth.Manager backed by the file store, printing the // authorization URL / device code to stdout. The store path honors -// ZERO_OAUTH_TOKENS_PATH (env), so callers/tests can redirect it. +// ZERO_OAUTH_TOKENS_PATH (env), so callers/tests can redirect it. Setting +// ZERO_OAUTH_STORAGE=encrypted-file selects the AES-256-GCM encrypted-at-rest +// backend (a per-user secret is created beside the token file). func newAuthManager(deps appDeps, out io.Writer) (*oauth.Manager, error) { - store, err := oauth.NewStore(oauth.StoreOptions{Now: deps.now}) + // Validate ZERO_OAUTH_STORAGE up front: a mistyped non-empty value must fail + // fast rather than silently fall back to plaintext while the user believes + // encryption is on. Empty = default (plaintext 0600); "encrypted-file" = AES. + encrypted := false + if mode := strings.TrimSpace(os.Getenv("ZERO_OAUTH_STORAGE")); mode != "" { + if !strings.EqualFold(mode, "encrypted-file") { + return nil, fmt.Errorf("invalid ZERO_OAUTH_STORAGE %q (supported: encrypted-file)", mode) + } + encrypted = true + } + store, err := oauth.NewStore(oauth.StoreOptions{ + Now: deps.now, + Encrypted: encrypted, + }) if err != nil { return nil, err } @@ -333,6 +349,10 @@ built in). For a provider named , set: ZERO_OAUTH__SCOPES ZERO_OAUTH__FLOW (loopback|device) Endpoint URLs must be https (loopback exempt). +Storage: tokens are written 0600 under $XDG_CONFIG_HOME/zero (override with +ZERO_OAUTH_TOKENS_PATH). Set ZERO_OAUTH_STORAGE=encrypted-file to encrypt them +at rest with AES-256-GCM (a per-user secret is created beside the token file). + Flags: --device Use the device-code flow (headless/SSH; no browser) --scope Add an OAuth scope (repeatable) diff --git a/internal/cli/auth_test.go b/internal/cli/auth_test.go index 235d51d65..1da39d516 100644 --- a/internal/cli/auth_test.go +++ b/internal/cli/auth_test.go @@ -17,6 +17,20 @@ func withAuthStore(t *testing.T) string { return path } +func TestRunAuthRejectsInvalidStorageMode(t *testing.T) { + withAuthStore(t) + // A mistyped value must fail fast, not silently fall back to plaintext while + // the user believes encryption is active. + t.Setenv("ZERO_OAUTH_STORAGE", "encryptd") + var stdout, stderr bytes.Buffer + if code := runWithDeps([]string{"auth", "status"}, &stdout, &stderr, appDeps{}); code == exitSuccess { + t.Fatalf("invalid ZERO_OAUTH_STORAGE should fail, got success; stdout=%q", stdout.String()) + } + if !strings.Contains(stderr.String(), "ZERO_OAUTH_STORAGE") { + t.Fatalf("error should name the offending env var, stderr=%q", stderr.String()) + } +} + func TestRunAuthStatusEmpty(t *testing.T) { withAuthStore(t) var stdout, stderr bytes.Buffer diff --git a/internal/oauth/encrypt.go b/internal/oauth/encrypt.go new file mode 100644 index 000000000..7a241cbd1 --- /dev/null +++ b/internal/oauth/encrypt.go @@ -0,0 +1,147 @@ +package oauth + +import ( + "crypto/aes" + "crypto/cipher" + "crypto/rand" + "errors" + "fmt" + "io" + "os" + "path/filepath" + "time" +) + +// secretBytes is the AES-256 key length kept in the per-user secret file. +const secretBytes = 32 + +// aesGCMCrypter encrypts the token file at rest with AES-256-GCM under a +// per-user random secret persisted (0600) beside the token file. The on-disk +// blob is nonce || ciphertext; GCM provides confidentiality AND tamper +// detection, so a corrupted/forged file fails closed on open. This is the +// opt-in "encrypted-file" storage backend; the default backend writes the +// 0600 plaintext JSON unchanged. +type aesGCMCrypter struct { + secretPath string +} + +func newAESGCMCrypter(secretPath string) *aesGCMCrypter { + return &aesGCMCrypter{secretPath: secretPath} +} + +// aead loads (or, when create is set, generates) the secret and returns the GCM +// AEAD. open passes create=false so a missing secret is a hard error rather than +// silently minting a new key that could never decrypt the existing file. +func (c *aesGCMCrypter) aead(create bool) (cipher.AEAD, error) { + secret, err := loadOrCreateSecret(c.secretPath, create) + if err != nil { + return nil, err + } + block, err := aes.NewCipher(secret) + if err != nil { + return nil, fmt.Errorf("oauth: build cipher: %w", err) + } + return cipher.NewGCM(block) +} + +// seal encrypts plaintext, prefixing a fresh random nonce. It creates the secret +// on first use. +func (c *aesGCMCrypter) seal(plaintext []byte) ([]byte, error) { + gcm, err := c.aead(true) + if err != nil { + return nil, err + } + nonce := make([]byte, gcm.NonceSize()) + if _, err := io.ReadFull(rand.Reader, nonce); err != nil { + return nil, fmt.Errorf("oauth: generate nonce: %w", err) + } + return gcm.Seal(nonce, nonce, plaintext, nil), nil +} + +// open decrypts a nonce||ciphertext blob, failing closed on a missing secret, +// a short blob, or a failed authentication tag (tampering / wrong key). +func (c *aesGCMCrypter) open(blob []byte) ([]byte, error) { + gcm, err := c.aead(false) + if err != nil { + return nil, err + } + if len(blob) < gcm.NonceSize() { + return nil, errors.New("oauth: encrypted token file is too short") + } + nonce, ciphertext := blob[:gcm.NonceSize()], blob[gcm.NonceSize():] + plaintext, err := gcm.Open(nil, nonce, ciphertext, nil) + if err != nil { + return nil, fmt.Errorf("oauth: decrypt token file (wrong secret or tampered): %w", err) + } + return plaintext, nil +} + +// loadOrCreateSecret reads the 32-byte secret at path. When create is set and +// the file is absent, it generates a random secret and creates the file +// exclusively (0600). A wrong-sized existing secret fails closed (corruption). +func loadOrCreateSecret(path string, create bool) ([]byte, error) { + if data, err := readSecretFile(path); err == nil { + return data, nil + } else if !errors.Is(err, os.ErrNotExist) { + return nil, err + } + if !create { + return nil, fmt.Errorf("oauth: token secret %s is missing; cannot decrypt the token file", path) + } + secret := make([]byte, secretBytes) + if _, err := io.ReadFull(rand.Reader, secret); err != nil { + return nil, fmt.Errorf("oauth: generate token secret: %w", err) + } + if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil { + return nil, err + } + // Create exclusively: an os.Rename publish would clobber on POSIX, so two + // concurrent first-run processes could each generate a secret and the loser + // would silently orphan the tokens it encrypts. O_EXCL lets exactly one + // process win; everyone else adopts the winner's on-disk secret. + f, err := os.OpenFile(path, os.O_CREATE|os.O_EXCL|os.O_WRONLY, 0o600) + if err != nil { + if errors.Is(err, os.ErrExist) { + // The winner may not have finished writing yet, so a read here can see a + // short/absent file. Retry briefly so concurrent first-run invocations + // converge on the winner's secret instead of failing transiently. + var lastErr error + for attempt := 0; attempt < 50; attempt++ { + secret, rerr := readSecretFile(path) + if rerr == nil { + return secret, nil + } + lastErr = rerr + time.Sleep(2 * time.Millisecond) + } + return nil, lastErr + } + return nil, fmt.Errorf("oauth: create token secret: %w", err) + } + if _, werr := f.Write(secret); werr != nil { + _ = f.Close() + _ = os.Remove(path) + return nil, fmt.Errorf("oauth: write token secret: %w", werr) + } + if cerr := f.Close(); cerr != nil { + _ = os.Remove(path) + return nil, fmt.Errorf("oauth: write token secret: %w", cerr) + } + return secret, nil +} + +// readSecretFile reads and validates the secret at path, returning a wrapped +// os.ErrNotExist when it is absent so callers can branch on creation. +func readSecretFile(path string) ([]byte, error) { + data, err := os.ReadFile(path) + if err != nil { + if errors.Is(err, os.ErrNotExist) { + return nil, err + } + return nil, fmt.Errorf("oauth: read token secret: %w", err) + } + if len(data) != secretBytes { + return nil, fmt.Errorf("oauth: token secret at %s has unexpected size %d", path, len(data)) + } + return data, nil +} diff --git a/internal/oauth/encrypt_test.go b/internal/oauth/encrypt_test.go new file mode 100644 index 000000000..3a4c6f4d4 --- /dev/null +++ b/internal/oauth/encrypt_test.go @@ -0,0 +1,180 @@ +package oauth + +import ( + "bytes" + "os" + "path/filepath" + "runtime" + "strings" + "sync" + "testing" +) + +func TestLoadOrCreateSecretConcurrentConverges(t *testing.T) { + path := filepath.Join(t.TempDir(), "tok.json.secret") + const n = 16 + secrets := make([][]byte, n) + errs := make([]error, n) + var wg sync.WaitGroup + for i := 0; i < n; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + secrets[i], errs[i] = loadOrCreateSecret(path, true) + }(i) + } + wg.Wait() + // Exactly one creator wins; every racer must converge on the same on-disk + // secret rather than orphaning its own (the O_EXCL + bounded-retry path). + for i := 0; i < n; i++ { + if errs[i] != nil { + t.Fatalf("goroutine %d: %v", i, errs[i]) + } + if !bytes.Equal(secrets[i], secrets[0]) { + t.Fatalf("goroutine %d got a divergent secret; concurrent create did not converge", i) + } + } +} + +func TestAESGCMCrypterRoundTripAndTamper(t *testing.T) { + secretPath := filepath.Join(t.TempDir(), "tok.json.secret") + c := newAESGCMCrypter(secretPath) + plaintext := []byte(`{"schemaVersion":1,"tokens":{}}`) + blob, err := c.seal(plaintext) + if err != nil { + t.Fatalf("seal: %v", err) + } + if bytes.Contains(blob, plaintext) { + t.Fatal("sealed blob must not contain the plaintext") + } + got, err := c.open(blob) + if err != nil { + t.Fatalf("open: %v", err) + } + if !bytes.Equal(got, plaintext) { + t.Fatalf("round-trip mismatch: %q", got) + } + // Tamper: flipping any byte breaks the GCM tag. + tampered := append([]byte(nil), blob...) + tampered[len(tampered)-1] ^= 0xff + if _, err := c.open(tampered); err == nil { + t.Fatal("open must reject a tampered blob (GCM auth)") + } +} + +func TestLoadOrCreateSecret(t *testing.T) { + path := filepath.Join(t.TempDir(), "tok.json.secret") + // Missing + create=false => fail closed (can't decrypt without the secret). + if _, err := loadOrCreateSecret(path, false); err == nil { + t.Fatal("missing secret with create=false must error") + } + secret, err := loadOrCreateSecret(path, true) + if err != nil { + t.Fatalf("create: %v", err) + } + if len(secret) != secretBytes { + t.Fatalf("secret length = %d, want %d", len(secret), secretBytes) + } + // Stable across reads. + again, err := loadOrCreateSecret(path, false) + if err != nil || !bytes.Equal(again, secret) { + t.Fatalf("secret not stable: %v", err) + } + if runtime.GOOS != "windows" { + info, err := os.Stat(path) + if err != nil { + t.Fatalf("stat secret: %v", err) + } + if perm := info.Mode().Perm(); perm != 0o600 { + t.Fatalf("secret file mode = %o, want 600", perm) + } + } + // Wrong-sized secret fails closed. + if err := os.WriteFile(path, []byte("short"), 0o600); err != nil { + t.Fatalf("corrupt secret: %v", err) + } + if _, err := loadOrCreateSecret(path, true); err == nil { + t.Fatal("wrong-sized secret must error") + } +} + +func newEncryptedStore(t *testing.T) (*Store, string) { + t.Helper() + path := filepath.Join(t.TempDir(), "oauth-tokens.json") + s, err := NewStore(StoreOptions{FilePath: path, Encrypted: true}) + if err != nil { + t.Fatalf("NewStore(encrypted): %v", err) + } + return s, path +} + +func TestEncryptedStoreRoundTripAndCiphertextOnDisk(t *testing.T) { + s, path := newEncryptedStore(t) + tok := Token{AccessToken: "super-secret-access", RefreshToken: "super-secret-refresh", Account: "me@x"} + if err := s.Save(ProviderKey("demo"), tok); err != nil { + t.Fatalf("Save: %v", err) + } + // On-disk file must be ciphertext: not valid JSON, no plaintext token. + raw, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read file: %v", err) + } + if strings.Contains(string(raw), "super-secret-access") || strings.Contains(string(raw), "schemaVersion") { + t.Fatalf("token file is not encrypted at rest:\n%s", raw) + } + // The secret file exists beside it. + if _, err := os.Stat(path + ".secret"); err != nil { + t.Fatalf("secret file missing: %v", err) + } + // A fresh Store (same path) decrypts and reads it back. + s2, err := NewStore(StoreOptions{FilePath: path, Encrypted: true}) + if err != nil { + t.Fatalf("NewStore 2: %v", err) + } + got, ok, err := s2.Load(ProviderKey("demo")) + if err != nil || !ok || got.AccessToken != "super-secret-access" || got.Account != "me@x" { + t.Fatalf("Load = %+v ok=%v err=%v", got, ok, err) + } + // Delete + Status work through the encrypted backend too. + statuses, err := s2.Status(KeyPrefixProvider) + if err != nil || len(statuses) != 1 { + t.Fatalf("Status = %+v err=%v", statuses, err) + } + if removed, err := s2.Delete(ProviderKey("demo")); err != nil || !removed { + t.Fatalf("Delete = %v %v", removed, err) + } +} + +func TestEncryptedStoreTamperFailsClosed(t *testing.T) { + s, path := newEncryptedStore(t) + if err := s.Save(ProviderKey("demo"), Token{AccessToken: "a"}); err != nil { + t.Fatalf("Save: %v", err) + } + raw, err := os.ReadFile(path) + if err != nil { + t.Fatalf("read file: %v", err) + } + if len(raw) == 0 { + t.Fatal("encrypted token file is empty") + } + raw[len(raw)-1] ^= 0xff + if err := os.WriteFile(path, raw, 0o600); err != nil { + t.Fatalf("rewrite: %v", err) + } + if _, _, err := s.Load(ProviderKey("demo")); err == nil { + t.Fatal("a tampered encrypted store must fail closed") + } +} + +func TestEncryptedStoreMissingSecretFailsClosed(t *testing.T) { + s, path := newEncryptedStore(t) + if err := s.Save(ProviderKey("demo"), Token{AccessToken: "a"}); err != nil { + t.Fatalf("Save: %v", err) + } + if err := os.Remove(path + ".secret"); err != nil { + t.Fatalf("remove secret: %v", err) + } + if _, _, err := s.Load(ProviderKey("demo")); err == nil { + t.Fatal("missing secret must fail closed, not return empty") + } +} diff --git a/internal/oauth/store.go b/internal/oauth/store.go index f9c36ded0..18ecc1255 100644 --- a/internal/oauth/store.go +++ b/internal/oauth/store.go @@ -54,12 +54,18 @@ type StoreOptions struct { FilePath string Env map[string]string Now func() time.Time + // Encrypted selects the AES-256-GCM encrypted-at-rest backend (a per-user + // secret is created beside the token file). Default (false) writes the 0600 + // plaintext JSON unchanged. + Encrypted bool } -// Store persists OAuth tokens (provider + MCP namespaces) in a 0600 JSON file, -// guarded by a cross-process lock and written atomically. +// Store persists OAuth tokens (provider + MCP namespaces) in a 0600 file, +// guarded by a cross-process lock and written atomically. The file is plaintext +// JSON by default, or AES-256-GCM ciphertext when the encrypted backend is on. type Store struct { filePath string + crypter *aesGCMCrypter // nil => plaintext backend now func() time.Time mu sync.Mutex } @@ -119,7 +125,11 @@ func NewStore(options StoreOptions) (*Store, error) { if now == nil { now = time.Now } - return &Store{filePath: filepath.Clean(filePath), now: now}, nil + store := &Store{filePath: filepath.Clean(filePath), now: now} + if options.Encrypted { + store.crypter = newAESGCMCrypter(store.filePath + ".secret") + } + return store, nil } // Save persists a token under key, replacing any existing entry. @@ -225,6 +235,12 @@ func (s *Store) readState() (storeFile, error) { } return storeFile{}, err } + if s.crypter != nil { + data, err = s.crypter.open(data) + if err != nil { + return storeFile{}, err + } + } var state storeFile if err := json.Unmarshal(data, &state); err != nil { return storeFile{}, fmt.Errorf("oauth: invalid token file at %s: %w", s.filePath, err) @@ -251,8 +267,16 @@ func (s *Store) writeState(state storeFile) error { if err != nil { return err } + payload := append(data, '\n') + if s.crypter != nil { + // Encrypted backend: the on-disk file is opaque ciphertext, not JSON. + payload, err = s.crypter.seal(data) + if err != nil { + return err + } + } tempPath := fmt.Sprintf("%s.tmp-%d-%d", s.filePath, os.Getpid(), s.now().UnixNano()) - if err := os.WriteFile(tempPath, append(data, '\n'), 0o600); err != nil { + if err := os.WriteFile(tempPath, payload, 0o600); err != nil { return err } if err := os.Rename(tempPath, s.filePath); err != nil {