diff --git a/go.mod b/go.mod index 7f588b5f1..c6cebf45d 100644 --- a/go.mod +++ b/go.mod @@ -10,6 +10,7 @@ require ( github.com/jackc/pgx/v5 v5.7.6 github.com/mark3labs/mcp-go v0.44.0 golang.org/x/net v0.52.0 + golang.org/x/sys v0.43.0 modernc.org/sqlite v1.45.0 ) @@ -59,7 +60,6 @@ require ( golang.org/x/exp v0.0.0-20251023183803-a4bb9ffd2546 // indirect golang.org/x/mod v0.34.0 // indirect golang.org/x/sync v0.20.0 // indirect - golang.org/x/sys v0.43.0 // indirect golang.org/x/text v0.36.0 // indirect golang.org/x/tools v0.43.0 // indirect gopkg.in/yaml.v3 v3.0.1 // indirect diff --git a/internal/store/generation_fence.go b/internal/store/generation_fence.go index 6b2bede89..f8f59c8cc 100644 --- a/internal/store/generation_fence.go +++ b/internal/store/generation_fence.go @@ -80,8 +80,15 @@ func ensureDatabaseFile(path string) error { } var openDB = func(dbPath string, generation *databaseGeneration) (*sql.DB, error) { - d := &generationDriver{Driver: &sqlite.Driver{}, generation: generation} - return sql.OpenDB(generationConnector{driver: d, name: dbPath}), nil + sqliteDriver := &sqlite.Driver{} + sqliteDriver.RegisterConnectionHook(func(conn sqlite.ExecQuerierContext, _ string) error { + if fc, ok := conn.(sqlite.FileControl); ok { + _, _ = fc.FileControlPersistWAL("main", 1) + } + return nil + }) + d := &generationDriver{Driver: sqliteDriver, generation: generation} + return sql.OpenDB(generationConnector{driver: d, name: storeDSN(dbPath)}), nil } type generationConnector struct { @@ -267,6 +274,28 @@ func (c generationConn) CheckNamedValue(value *driver.NamedValue) error { return driver.ErrSkip } +// FileControlPersistWAL forwards modernc's optional FileControl interface +// through the generation fence. database/sql exposes this wrapped connection +// to primeConnection via Conn.Raw, so omitting it would make persistent WAL +// unavailable whenever generation fencing is enabled. +func (c generationConn) FileControlPersistWAL(dbName string, mode int) (int, error) { + if err := c.generation.check(); err != nil { + return 0, err + } + fc, ok := c.Conn.(sqlite.FileControl) + if !ok { + return 0, errors.New("database connection does not implement sqlite.FileControl") + } + result, err := fc.FileControlPersistWAL(dbName, mode) + if err != nil { + return 0, err + } + if err := c.generation.check(); err != nil { + return 0, err + } + return result, nil +} + type generationStmt struct { driver.Stmt generation *databaseGeneration diff --git a/internal/store/generation_fence_test.go b/internal/store/generation_fence_test.go index 36a9f939a..dbcf4f1ab 100644 --- a/internal/store/generation_fence_test.go +++ b/internal/store/generation_fence_test.go @@ -8,6 +8,8 @@ import ( "os" "path/filepath" "testing" + + sqlite "modernc.org/sqlite" ) func TestDatabaseGeneration(t *testing.T) { @@ -110,6 +112,24 @@ func TestGenerationFenceRejectsUnsafeOperations(t *testing.T) { }) } +func TestGenerationConnExposesFileControlForPrimeConnection(t *testing.T) { + generation, _ := newTestDatabaseGeneration(t, false, false) + base := &testFenceFileControlConn{} + conn := generationConn{Conn: base, generation: generation} + + fc, ok := any(conn).(sqlite.FileControl) + if !ok { + t.Fatal("generation connection does not expose sqlite.FileControl") + } + mode, err := fc.FileControlPersistWAL("main", 1) + if err != nil { + t.Fatalf("FileControlPersistWAL: %v", err) + } + if mode != 1 || base.dbName != "main" || base.mode != 1 { + t.Fatalf("persist WAL = (%d, %q, %d), want (1, main, 1)", mode, base.dbName, base.mode) + } +} + func TestNewRejectsGenerationChangedBeforeSQLiteOpens(t *testing.T) { original := openDB t.Cleanup(func() { openDB = original }) @@ -209,6 +229,17 @@ type testFenceConn struct { rows *testFenceRows } +type testFenceFileControlConn struct { + testFenceConn + dbName string + mode int +} + +func (c *testFenceFileControlConn) FileControlPersistWAL(dbName string, mode int) (int, error) { + c.dbName, c.mode = dbName, mode + return mode, nil +} + func (c *testFenceConn) Prepare(string) (driver.Stmt, error) { return testFenceStmt{}, nil } func (c *testFenceConn) Close() error { return nil } func (c *testFenceConn) Begin() (driver.Tx, error) { return testFenceTx{}, nil } diff --git a/internal/store/migration_lock.go b/internal/store/migration_lock.go new file mode 100644 index 000000000..b1944f76b --- /dev/null +++ b/internal/store/migration_lock.go @@ -0,0 +1,68 @@ +package store + +import ( + "fmt" + "os" + "time" +) + +// migrationLockTimeout bounds how long a process waits for the migration +// lock. A hung holder must produce a loud, actionable error instead of +// silently blocking every engram process on the machine forever. It is a +// variable (not a constant) so tests can shorten the timeout path. +var migrationLockTimeout = 60 * time.Second + +// acquireMigrationLock takes an exclusive advisory lock on path and returns +// a function that releases it. It serializes whole processes around the +// migration suite and the startup repair so that the destructive +// check-then-act rebuilds inside migrate() can never run twice concurrently +// against the same database. +// +// Acquisition is non-blocking with a bounded growing backoff (up to +// migrationLockTimeout total) rather than a blocking lock: a stuck holder +// then surfaces as a clear error naming the lock file instead of a silent +// machine-wide hang. +// +// The lock file is deliberately left in place after unlock: unlinking it +// would open a race where a third process re-creates the path and locks a +// different inode/file object, defeating the exclusion. +func acquireMigrationLock(path string) (func(), error) { + f, err := os.OpenFile(path, os.O_CREATE|os.O_RDWR, 0o644) + if err != nil { + return nil, fmt.Errorf("open migration lock file %s: %w", path, err) + } + + deadline := time.Now().Add(migrationLockTimeout) + backoff := 10 * time.Millisecond + for { + acquired, err := tryLockMigrationFile(f) + if err != nil { + _ = f.Close() + return nil, fmt.Errorf("lock migration lock file %s: %w", path, err) + } + if acquired { + return func() { + _ = unlockMigrationFile(f) + _ = f.Close() + }, nil + } + if time.Now().After(deadline) { + _ = f.Close() + return nil, fmt.Errorf( + "timed out after %s waiting for migration lock %s — another engram process appears to be holding it; check for a stuck engram process (and terminate it) before retrying", + migrationLockTimeout, path, + ) + } + time.Sleep(backoff) + // Grow the poll interval but cap it low: healthy holders release the + // lock within milliseconds (the startup repair fast path is read-only), + // and every engram subcommand acquires this lock, so an aggressive cap + // keeps contended cold starts snappy. + if backoff < 100*time.Millisecond { + backoff *= 2 + if backoff > 100*time.Millisecond { + backoff = 100 * time.Millisecond + } + } + } +} diff --git a/internal/store/migration_lock_unix.go b/internal/store/migration_lock_unix.go new file mode 100644 index 000000000..fd777a862 --- /dev/null +++ b/internal/store/migration_lock_unix.go @@ -0,0 +1,30 @@ +//go:build aix || darwin || dragonfly || freebsd || linux || netbsd || openbsd + +package store + +import ( + "errors" + "os" + "syscall" +) + +// tryLockMigrationFile attempts a non-blocking exclusive flock(2) on f. +// It reports (false, nil) when another process (or file description) holds +// the lock, so the caller can retry with backoff. +func tryLockMigrationFile(f *os.File) (bool, error) { + err := syscall.Flock(int(f.Fd()), syscall.LOCK_EX|syscall.LOCK_NB) + if err == nil { + return true, nil + } + // EWOULDBLOCK/EAGAIN: lock is held elsewhere. EINTR: interrupted by a + // signal. Both are retryable, not failures. + if errors.Is(err, syscall.EWOULDBLOCK) || errors.Is(err, syscall.EAGAIN) || errors.Is(err, syscall.EINTR) { + return false, nil + } + return false, err +} + +// unlockMigrationFile releases the flock taken by tryLockMigrationFile. +func unlockMigrationFile(f *os.File) error { + return syscall.Flock(int(f.Fd()), syscall.LOCK_UN) +} diff --git a/internal/store/migration_lock_windows.go b/internal/store/migration_lock_windows.go new file mode 100644 index 000000000..92372024c --- /dev/null +++ b/internal/store/migration_lock_windows.go @@ -0,0 +1,34 @@ +//go:build windows + +package store + +import ( + "errors" + "os" + + "golang.org/x/sys/windows" +) + +// tryLockMigrationFile attempts a non-blocking exclusive LockFileEx on f. +// It reports (false, nil) when another process holds the lock, so the +// caller can retry with backoff. +func tryLockMigrationFile(f *os.File) (bool, error) { + ol := new(windows.Overlapped) + err := windows.LockFileEx( + windows.Handle(f.Fd()), + windows.LOCKFILE_EXCLUSIVE_LOCK|windows.LOCKFILE_FAIL_IMMEDIATELY, + 0, 1, 0, ol, + ) + if err == nil { + return true, nil + } + if errors.Is(err, windows.ERROR_LOCK_VIOLATION) { + return false, nil + } + return false, err +} + +// unlockMigrationFile releases the lock taken by tryLockMigrationFile. +func unlockMigrationFile(f *os.File) error { + return windows.UnlockFileEx(windows.Handle(f.Fd()), 0, 1, 0, new(windows.Overlapped)) +} diff --git a/internal/store/startup_gate_test.go b/internal/store/startup_gate_test.go new file mode 100644 index 000000000..3a44f6698 --- /dev/null +++ b/internal/store/startup_gate_test.go @@ -0,0 +1,577 @@ +package store + +// Tests for the SQLite-corruption fixes: +// +// 1. persistent WAL — closing the store must NOT unlink the -wal file +// (mechanism behind upstream #477/#571), +// 2. cold-start concurrency — many stores opening the same fresh database +// simultaneously, in-process and across child processes (#559), and +// 3. the user_version migration gate — the migration suite runs exactly +// once per schema generation and is never run against a database +// stamped by a newer engram. + +import ( + "context" + "database/sql" + "fmt" + "net/url" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "testing" + "time" +) + +func TestPersistentWALSurvivesClose(t *testing.T) { + cfg := mustDefaultConfig(t) + cfg.DataDir = t.TempDir() + + s, err := New(cfg) + if err != nil { + t.Fatalf("New: %v", err) + } + if err := s.CreateSession("wal-session", "wal-project", cfg.DataDir); err != nil { + t.Fatalf("CreateSession: %v", err) + } + + walPath := filepath.Join(cfg.DataDir, "engram.db-wal") + if _, err := os.Stat(walPath); err != nil { + t.Fatalf("-wal file missing while store is open: %v", err) + } + + if err := s.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + if _, err := os.Stat(walPath); err != nil { + t.Fatalf("-wal file was unlinked on close — persistent WAL is not active: %v", err) + } +} + +func TestNewWithQuestionMarkInDataDirectory(t *testing.T) { + if filepath.Separator != '/' { + t.Skip("Unix filesystem path behavior") + } + + cfg := mustDefaultConfig(t) + cfg.DataDir = filepath.Join(t.TempDir(), "data?query") + + s, err := New(cfg) + if err != nil { + t.Fatalf("New: %v", err) + } + + if err := s.CreateSession("question-mark-session", "question-mark-project", cfg.DataDir); err != nil { + t.Fatalf("CreateSession: %v", err) + } + if err := s.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + raw, err := sql.Open("sqlite", (&url.URL{Scheme: "file", Path: filepath.Join(cfg.DataDir, "engram.db")}).String()) + if err != nil { + t.Fatalf("open configured database: %v", err) + } + defer raw.Close() + + var sessions int + if err := raw.QueryRow("SELECT COUNT(*) FROM sessions").Scan(&sessions); err != nil { + t.Fatalf("count sessions in configured database: %v", err) + } + if sessions != 1 { + t.Fatalf("sessions in configured database = %d, want 1", sessions) + } +} + +func TestUserVersionGateSkipsSecondOpen(t *testing.T) { + cfg := mustDefaultConfig(t) + cfg.DataDir = t.TempDir() + + before := migrateRunCount.Load() + + s1, err := New(cfg) + if err != nil { + t.Fatalf("first New: %v", err) + } + var v int + if err := s1.db.QueryRow("PRAGMA user_version").Scan(&v); err != nil { + t.Fatalf("read user_version: %v", err) + } + if v != schemaVersion { + t.Fatalf("user_version after first open = %d, want %d", v, schemaVersion) + } + if err := s1.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + afterFirst := migrateRunCount.Load() + if got := afterFirst - before; got != 1 { + t.Fatalf("migration suite ran %d times on a fresh database, want exactly 1", got) + } + + s2, err := New(cfg) + if err != nil { + t.Fatalf("second New: %v", err) + } + defer s2.Close() + + if got := migrateRunCount.Load() - afterFirst; got != 0 { + t.Fatalf("migration suite ran %d times on an already-migrated database, want 0", got) + } + + // The gated (skipped-migration) store must still be fully usable. + if err := s2.CreateSession("gate-session", "gate-project", cfg.DataDir); err != nil { + t.Fatalf("CreateSession on gated store: %v", err) + } +} + +func TestNewerSchemaVersionIsLeftUntouched(t *testing.T) { + cfg := mustDefaultConfig(t) + cfg.DataDir = t.TempDir() + + s1, err := New(cfg) + if err != nil { + t.Fatalf("New: %v", err) + } + if err := s1.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + // Simulate a database already migrated by a NEWER engram. + future := schemaVersion + 7 + raw, err := sql.Open("sqlite", filepath.Join(cfg.DataDir, "engram.db")) + if err != nil { + t.Fatalf("raw open: %v", err) + } + if _, err := raw.Exec(fmt.Sprintf("PRAGMA user_version = %d", future)); err != nil { + t.Fatalf("stamp future user_version: %v", err) + } + if err := raw.Close(); err != nil { + t.Fatalf("raw close: %v", err) + } + + before := migrateRunCount.Load() + + s2, err := New(cfg) + if err != nil { + t.Fatalf("New on newer-schema database: %v", err) + } + defer s2.Close() + + if got := migrateRunCount.Load() - before; got != 0 { + t.Fatalf("migration suite ran %d times against a newer schema, want 0", got) + } + + var v int + if err := s2.db.QueryRow("PRAGMA user_version").Scan(&v); err != nil { + t.Fatalf("read user_version: %v", err) + } + if v != future { + t.Fatalf("user_version was rewritten to %d, want it left at %d", v, future) + } +} + +func TestNewerSchemaVersionSkipsSchemaRepair(t *testing.T) { + cfg := mustDefaultConfig(t) + cfg.DataDir = t.TempDir() + + s1, err := New(cfg) + if err != nil { + t.Fatalf("New: %v", err) + } + if err := s1.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + raw, err := sql.Open("sqlite", filepath.Join(cfg.DataDir, "engram.db")) + if err != nil { + t.Fatalf("raw open: %v", err) + } + if _, err := raw.Exec(` + INSERT INTO sync_apply_deferred (sync_id, entity, payload) + VALUES ('future-schema-repair', 'relation', '{"sync_id":"derived-by-migrate"}') + `); err != nil { + t.Fatalf("seed schema repair row: %v", err) + } + if _, err := raw.Exec(fmt.Sprintf("PRAGMA user_version = %d", schemaVersion+7)); err != nil { + t.Fatalf("stamp future user_version: %v", err) + } + if err := raw.Close(); err != nil { + t.Fatalf("close seeded database: %v", err) + } + + s2, err := New(cfg) + if err != nil { + t.Fatalf("New on newer-schema database: %v", err) + } + defer s2.Close() + + var payloadSyncID string + if err := s2.db.QueryRow(`SELECT payload_sync_id FROM sync_apply_deferred WHERE sync_id = 'future-schema-repair'`).Scan(&payloadSyncID); err != nil { + t.Fatalf("read schema repair row: %v", err) + } + if payloadSyncID != "" { + t.Fatalf("future-schema open derived payload_sync_id %q, want no schema repair write", payloadSyncID) + } +} + +func TestConcurrentColdStartGoroutines(t *testing.T) { + cfg := mustDefaultConfig(t) + cfg.DataDir = t.TempDir() + + before := migrateRunCount.Load() + + const n = 8 + start := make(chan struct{}) + errs := make(chan error, n) + var wg sync.WaitGroup + for i := 0; i < n; i++ { + wg.Add(1) + go func(i int) { + defer wg.Done() + <-start + s, err := New(cfg) + if err != nil { + errs <- fmt.Errorf("goroutine %d: New: %w", i, err) + return + } + defer s.Close() + if err := s.CreateSession(fmt.Sprintf("cold-%d", i), "coldstart", cfg.DataDir); err != nil { + errs <- fmt.Errorf("goroutine %d: CreateSession: %w", i, err) + return + } + errs <- nil + }(i) + } + close(start) + wg.Wait() + close(errs) + for err := range errs { + if err != nil { + t.Error(err) + } + } + if t.Failed() { + return + } + + if got := migrateRunCount.Load() - before; got != 1 { + t.Errorf("migration suite ran %d times across %d concurrent cold starts, want exactly 1", got, n) + } + + // Verify the resulting database is complete and healthy. + s, err := New(cfg) + if err != nil { + t.Fatalf("verification New: %v", err) + } + defer s.Close() + + var sessions int + if err := s.db.QueryRow("SELECT COUNT(*) FROM sessions").Scan(&sessions); err != nil { + t.Fatalf("count sessions: %v", err) + } + if sessions != n { + t.Errorf("sessions = %d, want %d", sessions, n) + } + var v int + if err := s.db.QueryRow("PRAGMA user_version").Scan(&v); err != nil { + t.Fatalf("read user_version: %v", err) + } + if v != schemaVersion { + t.Errorf("user_version = %d, want %d", v, schemaVersion) + } + var integrity string + if err := s.db.QueryRow("PRAGMA integrity_check").Scan(&integrity); err != nil { + t.Fatalf("integrity_check: %v", err) + } + if integrity != "ok" { + t.Errorf("integrity_check = %q, want ok", integrity) + } +} + +// coldStartChildEnv points a re-executed child copy of the test binary at the +// shared data directory used by TestConcurrentColdStartProcesses. +const coldStartChildEnv = "ENGRAM_TEST_COLDSTART_DIR" + +// TestColdStartChildProcess is not a standalone test: it is re-executed as a +// child process by TestConcurrentColdStartProcesses and skips otherwise. +func TestColdStartChildProcess(t *testing.T) { + dir := os.Getenv(coldStartChildEnv) + if dir == "" { + t.Skip("runs only as a child of TestConcurrentColdStartProcesses") + } + + cfg, err := DefaultConfig() + if err != nil { + t.Fatalf("DefaultConfig: %v", err) + } + cfg.DataDir = dir + cfg.DedupeWindow = time.Hour + + s, err := New(cfg) + if err != nil { + t.Fatalf("child cold-start New: %v", err) + } + defer s.Close() + + if err := s.CreateSession(fmt.Sprintf("proc-%d", os.Getpid()), "coldstart-proc", dir); err != nil { + t.Fatalf("child CreateSession: %v", err) + } +} + +func TestConcurrentColdStartProcesses(t *testing.T) { + if os.Getenv(coldStartChildEnv) != "" { + t.Skip("child process mode") + } + + exe, err := os.Executable() + if err != nil { + t.Fatalf("os.Executable: %v", err) + } + + dir := t.TempDir() + const procs = 3 + + cmds := make([]*exec.Cmd, procs) + outputs := make([]*strings.Builder, procs) + for i := range cmds { + outputs[i] = &strings.Builder{} + cmd := exec.Command(exe, "-test.run", "^TestColdStartChildProcess$", "-test.v", "-test.timeout", "60s") + cmd.Env = append(os.Environ(), coldStartChildEnv+"="+dir) + cmd.Stdout = outputs[i] + cmd.Stderr = outputs[i] + if err := cmd.Start(); err != nil { + t.Fatalf("start child %d: %v", i, err) + } + cmds[i] = cmd + } + for i, cmd := range cmds { + if err := cmd.Wait(); err != nil { + t.Errorf("child process %d failed: %v\noutput:\n%s", i, err, outputs[i].String()) + } + } + if t.Failed() { + return + } + + // Every child cold-started against the same fresh database and wrote one + // session. Open from the parent and verify the result is consistent. + cfg, err := DefaultConfig() + if err != nil { + t.Fatalf("DefaultConfig: %v", err) + } + cfg.DataDir = dir + cfg.DedupeWindow = time.Hour + + s, err := New(cfg) + if err != nil { + t.Fatalf("parent verification New: %v", err) + } + defer s.Close() + + var sessions int + if err := s.db.QueryRow("SELECT COUNT(*) FROM sessions").Scan(&sessions); err != nil { + t.Fatalf("count sessions: %v", err) + } + if sessions != procs { + t.Errorf("sessions = %d, want %d", sessions, procs) + } + var v int + if err := s.db.QueryRow("PRAGMA user_version").Scan(&v); err != nil { + t.Fatalf("read user_version: %v", err) + } + if v != schemaVersion { + t.Errorf("user_version = %d, want %d", v, schemaVersion) + } + var integrity string + if err := s.db.QueryRow("PRAGMA integrity_check").Scan(&integrity); err != nil { + t.Fatalf("integrity_check: %v", err) + } + if integrity != "ok" { + t.Errorf("integrity_check = %q, want ok", integrity) + } +} + +// TestConnectionReplacementKeepsConfiguration guards the pool-replacement +// regression: database/sql silently discards a modernc connection after a +// context-cancelled query interrupts it (IsValid/ResetSession fail once +// sqlite3_is_interrupted) and opens a fresh one. Because the pragmas travel +// in the DSN and persist-WAL is applied by a driver connection hook, the +// replacement connection must come up fully configured — busy_timeout 5000 +// (#559) and persistent WAL (#477) intact. +func TestConnectionReplacementKeepsConfiguration(t *testing.T) { + cfg := mustDefaultConfig(t) + cfg.DataDir = t.TempDir() + + s, err := New(cfg) + if err != nil { + t.Fatalf("New: %v", err) + } + defer s.Close() + + // Plant a session-scoped marker on the current physical connection. + // PRAGMA cache_size is per-connection and not part of the DSN, so its + // disappearance later proves the pool swapped in a new connection. + if _, err := s.db.Exec("PRAGMA cache_size = -12345"); err != nil { + t.Fatalf("set cache_size marker: %v", err) + } + var marker int + if err := s.db.QueryRow("PRAGMA cache_size").Scan(&marker); err != nil { + t.Fatalf("read cache_size marker: %v", err) + } + if marker != -12345 { + t.Fatalf("cache_size marker = %d, want -12345", marker) + } + + // Interrupt a long-running query via context timeout. modernc's + // interruptOnDone calls sqlite3_interrupt, poisoning the connection so + // the pool discards it on release. + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + var n int64 + if err := s.db.QueryRowContext(ctx, + `WITH RECURSIVE c(x) AS (SELECT 1 UNION ALL SELECT x+1 FROM c) SELECT count(*) FROM c`, + ).Scan(&n); err == nil { + t.Fatal("expected the interrupted query to fail, it succeeded") + } + + // The pool must have replaced the physical connection... + var cacheSize int + if err := s.db.QueryRow("PRAGMA cache_size").Scan(&cacheSize); err != nil { + t.Fatalf("query after interruption: %v", err) + } + if cacheSize == -12345 { + t.Fatal("cache_size marker survived — connection was not replaced; test cannot exercise the replacement path") + } + + // ...and the replacement must be fully configured. + var busy int + if err := s.db.QueryRow("PRAGMA busy_timeout").Scan(&busy); err != nil { + t.Fatalf("read busy_timeout: %v", err) + } + if busy != 5000 { + t.Errorf("busy_timeout on replacement connection = %d, want 5000 (#559 regression)", busy) + } + var journalMode string + if err := s.db.QueryRow("PRAGMA journal_mode").Scan(&journalMode); err != nil { + t.Fatalf("read journal_mode: %v", err) + } + if !strings.EqualFold(journalMode, "wal") { + t.Errorf("journal_mode on replacement connection = %q, want wal", journalMode) + } + var foreignKeys int + if err := s.db.QueryRow("PRAGMA foreign_keys").Scan(&foreignKeys); err != nil { + t.Fatalf("read foreign_keys: %v", err) + } + if foreignKeys != 1 { + t.Errorf("foreign_keys on replacement connection = %d, want 1", foreignKeys) + } + + // Persist-WAL must be held on the replacement connection (query mode -1). + conn, err := s.db.Conn(context.Background()) + if err != nil { + t.Fatalf("pin replacement connection: %v", err) + } + if err := conn.Raw(func(driverConn any) error { + fc, ok := driverConn.(interface { + FileControlPersistWAL(string, int) (int, error) + }) + if !ok { + return fmt.Errorf("driver connection %T has no FileControlPersistWAL", driverConn) + } + mode, err := fc.FileControlPersistWAL("main", -1) + if err != nil { + return err + } + if mode != 1 { + return fmt.Errorf("persist-WAL mode on replacement connection = %d, want 1 (#477 regression)", mode) + } + return nil + }); err != nil { + t.Error(err) + } + if err := conn.Close(); err != nil { + t.Fatalf("release pinned connection: %v", err) + } + + // End to end: write through the replacement connection, close, and the + // -wal file must survive. + if err := s.CreateSession("replacement-session", "replacement-project", cfg.DataDir); err != nil { + t.Fatalf("CreateSession on replacement connection: %v", err) + } + if err := s.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + walPath := filepath.Join(cfg.DataDir, "engram.db-wal") + if _, err := os.Stat(walPath); err != nil { + t.Fatalf("-wal file was unlinked on close after connection replacement — persist-WAL was lost: %v", err) + } +} + +func TestAcquireMigrationLockTimesOutWithDiagnostic(t *testing.T) { + origTimeout := migrationLockTimeout + migrationLockTimeout = 300 * time.Millisecond + t.Cleanup(func() { migrationLockTimeout = origTimeout }) + + path := filepath.Join(t.TempDir(), ".migrate.lock") + + unlock, err := acquireMigrationLock(path) + if err != nil { + t.Fatalf("first acquire: %v", err) + } + defer unlock() + + start := time.Now() + _, err = acquireMigrationLock(path) + if err == nil { + t.Fatal("second acquire succeeded while the first lock was held; want timeout error") + } + if elapsed := time.Since(start); elapsed < 300*time.Millisecond { + t.Errorf("timed out after %s, want at least the %s budget", elapsed, 300*time.Millisecond) + } + if !strings.Contains(err.Error(), path) { + t.Errorf("timeout error does not name the lock file: %v", err) + } + if !strings.Contains(err.Error(), "stuck engram process") { + t.Errorf("timeout error does not point at a stuck process: %v", err) + } +} + +func TestAcquireMigrationLockExcludes(t *testing.T) { + path := filepath.Join(t.TempDir(), ".migrate.lock") + + unlock1, err := acquireMigrationLock(path) + if err != nil { + t.Fatalf("first acquire: %v", err) + } + + acquired := make(chan struct{}) + go func() { + unlock2, err := acquireMigrationLock(path) + if err != nil { + t.Errorf("second acquire: %v", err) + close(acquired) + return + } + close(acquired) + unlock2() + }() + + select { + case <-acquired: + t.Fatal("second acquire succeeded while the first lock was still held") + case <-time.After(150 * time.Millisecond): + // Still blocked — expected. + } + + unlock1() + + select { + case <-acquired: + // Granted after release — expected. + case <-time.After(5 * time.Second): + t.Fatal("second acquire did not proceed after the first lock was released") + } +} diff --git a/internal/store/store.go b/internal/store/store.go index 2437f89e9..03e257d30 100644 --- a/internal/store/store.go +++ b/internal/store/store.go @@ -15,6 +15,7 @@ import ( "errors" "fmt" "log" + "net/url" "os" "path/filepath" "regexp" @@ -23,6 +24,7 @@ import ( "strconv" "strings" "sync" + "sync/atomic" "time" "unicode" "unicode/utf8" @@ -707,6 +709,15 @@ func (s *Store) commitHook(tx *sql.Tx) error { } func New(cfg Config) (*Store, error) { + return newStore(cfg) +} + +// newWithoutRepair is retained as a test helper alias for New. +func newWithoutRepair(cfg Config) (*Store, error) { + return newStore(cfg) +} + +func newStore(cfg Config) (*Store, error) { if !filepath.IsAbs(cfg.DataDir) { return nil, fmt.Errorf("engram: data directory must be an absolute path, got %q — set ENGRAM_DATA_DIR or ensure your home directory is resolvable", cfg.DataDir) } @@ -734,78 +745,152 @@ func New(cfg Config) (*Store, error) { } }() - // SQLite performance pragmas - pragmas := []string{ - "PRAGMA journal_mode = WAL", - "PRAGMA busy_timeout = 5000", - "PRAGMA synchronous = NORMAL", - "PRAGMA foreign_keys = ON", - } - for _, p := range pragmas { - if _, err := db.Exec(p); err != nil { - return nil, fmt.Errorf("engram: pragma %q: %w", p, err) - } + if err := primeConnection(db); err != nil { + return nil, err } s := &Store{db: db, cfg: cfg, hooks: defaultStoreHooks()} - if err := s.migrate(); err != nil { - return nil, fmt.Errorf("engram: migration: %w", err) + if err := s.runStartupMigrations(); err != nil { + return nil, err } succeeded = true return s, nil } -// newWithoutRepair is retained as a test helper alias for New. Enrolled-project -// repair is deferred until the first synchronization operation in both cases. -func newWithoutRepair(cfg Config) (*Store, error) { - if !filepath.IsAbs(cfg.DataDir) { - return nil, fmt.Errorf("engram: data directory must be an absolute path, got %q — set ENGRAM_DATA_DIR or ensure your home directory is resolvable", cfg.DataDir) +func (s *Store) Close() error { + return s.db.Close() +} + +// ─── Migrations ────────────────────────────────────────────────────────────── + +// storeDSN builds the modernc.org/sqlite DSN for dbPath with the session +// pragmas encoded as _pragma query parameters. The driver applies these to +// every physical connection it opens, including replacement pool connections. +func storeDSN(dbPath string) string { + q := url.Values{} + for _, p := range []string{ + "busy_timeout(5000)", + "journal_mode(WAL)", + "synchronous(NORMAL)", + "foreign_keys(1)", + } { + q.Add("_pragma", p) } - if err := os.MkdirAll(cfg.DataDir, 0755); err != nil { - return nil, fmt.Errorf("engram: create data dir: %w", err) + if filepath.Separator == '/' { + return (&url.URL{Scheme: "file", Path: dbPath}).String() + "?" + q.Encode() } + return dbPath + "?" + q.Encode() +} - dbPath := filepath.Join(cfg.DataDir, "engram.db") - if err := ensureDatabaseFile(dbPath); err != nil { - return nil, fmt.Errorf("engram: create database file: %w", err) +var walSwitchRetryBackoffs = []time.Duration{ + 10 * time.Millisecond, + 25 * time.Millisecond, + 50 * time.Millisecond, + 100 * time.Millisecond, + 200 * time.Millisecond, +} + +// primeConnection opens the first physical connection with a bounded retry +// while concurrent cold starts race the rollback-journal-to-WAL conversion. +func primeConnection(db *sql.DB) error { + ctx := context.Background() + var conn *sql.Conn + var lastErr error + for attempt := 0; attempt <= len(walSwitchRetryBackoffs); attempt++ { + conn, lastErr = db.Conn(ctx) + if lastErr == nil { + break + } + if !isRetryableSQLiteLockError(lastErr) || attempt == len(walSwitchRetryBackoffs) { + return fmt.Errorf("engram: open initial connection: %w", lastErr) + } + time.Sleep(walSwitchRetryBackoffs[attempt]) } - generation, err := newDatabaseGeneration(dbPath) + defer conn.Close() + + return conn.Raw(func(driverConn any) error { + fc, ok := driverConn.(sqlite.FileControl) + if !ok { + return fmt.Errorf("engram: driver connection %T does not implement sqlite.FileControl", driverConn) + } + mode, err := fc.FileControlPersistWAL("main", 1) + if err != nil { + return fmt.Errorf("engram: enable persistent WAL: %w", err) + } + if mode != 1 { + return fmt.Errorf("engram: persistent WAL not active (file control returned mode %d)", mode) + } + return nil + }) +} + +const schemaVersion = 1 + +var migrateRunCount atomic.Int64 + +// runStartupMigrations serializes migrations and the every-open repair. It +// never writes a database with a schema version newer than this binary. +func (s *Store) runStartupMigrations() error { + current, err := s.readUserVersion() if err != nil { - return nil, fmt.Errorf("engram: capture database generation: %w", err) + return fmt.Errorf("engram: read user_version: %w", err) } - db, err := openDB(dbPath, generation) + if isFutureSchemaVersion(current) { + return nil + } + unlock, err := acquireMigrationLock(filepath.Join(s.cfg.DataDir, ".migrate.lock")) if err != nil { - return nil, fmt.Errorf("engram: open database: %w", err) + return fmt.Errorf("engram: acquire migration lock: %w", err) } - db.SetMaxOpenConns(1) + defer unlock() - pragmas := []string{ - "PRAGMA journal_mode = WAL", - "PRAGMA busy_timeout = 5000", - "PRAGMA synchronous = NORMAL", - "PRAGMA foreign_keys = ON", + // Re-read under the lock before deciding whether any startup write is safe: + // another process may have migrated this database, or a newer binary may + // have replaced it, while this process waited for the lock. + current, err = s.readUserVersion() + if err != nil { + return fmt.Errorf("engram: re-read user_version: %w", err) } - for _, p := range pragmas { - if _, err := db.Exec(p); err != nil { - _ = db.Close() - return nil, fmt.Errorf("engram: pragma %q: %w", p, err) - } + if isFutureSchemaVersion(current) { + return nil } - s := &Store{db: db, cfg: cfg, hooks: defaultStoreHooks()} + needsVersionStamp := current < schemaVersion + if needsVersionStamp { + migrateRunCount.Add(1) + } if err := s.migrate(); err != nil { - _ = db.Close() - return nil, fmt.Errorf("engram: migration: %w", err) + return fmt.Errorf("engram: migration: %w", err) } - return s, nil + if needsVersionStamp { + if err := s.setUserVersion(schemaVersion); err != nil { + return fmt.Errorf("engram: set user_version: %w", err) + } + } + return nil } -func (s *Store) Close() error { - return s.db.Close() +func isFutureSchemaVersion(current int) bool { + if current > schemaVersion { + log.Printf("[store] database schema version %d is newer than this binary's %d — skipping migrations (upgrade engram to manage this database)", current, schemaVersion) + return true + } + return false } -// ─── Migrations ────────────────────────────────────────────────────────────── +func (s *Store) readUserVersion() (int, error) { + var v int + if err := s.db.QueryRow("PRAGMA user_version").Scan(&v); err != nil { + return 0, err + } + return v, nil +} + +func (s *Store) setUserVersion(v int) error { + _, err := s.execHook(s.db, fmt.Sprintf("PRAGMA user_version = %d", v)) + return err +} func (s *Store) migrate() error { schema := `