diff --git a/internal/storage/backup.go b/internal/storage/backup.go index 0e81db1..c16f74b 100644 --- a/internal/storage/backup.go +++ b/internal/storage/backup.go @@ -42,6 +42,11 @@ func GetBackupDir() (string, error) { // CreateBackup uses SQLite's VACUUM INTO to produce a clean, consistent snapshot of the live database. func (d *Database) CreateBackup() (*BackupInfo, error) { + return d.createBackup("") +} + +// createBackup writes a VACUUM INTO snapshot to the backups directory. +func (d *Database) createBackup(suffix string) (*BackupInfo, error) { backupDir, err := GetBackupDir() if err != nil { return nil, err @@ -52,7 +57,7 @@ func (d *Database) CreateBackup() (*BackupInfo, error) { } now := time.Now() - filename := fmt.Sprintf("tmpo-%s.db", now.Format("20060102-150405")) + filename := fmt.Sprintf("tmpo-%s%s.db", now.Format("20060102-150405"), suffix) destPath := filepath.Join(backupDir, filename) escapedPath := strings.ReplaceAll(destPath, "'", "''") diff --git a/internal/storage/backup_test.go b/internal/storage/backup_test.go index a49f9dc..812be5c 100644 --- a/internal/storage/backup_test.go +++ b/internal/storage/backup_test.go @@ -328,6 +328,138 @@ func TestRestoreBackup_SucceedsWithoutSidecars(t *testing.T) { assert.NoError(t, RestoreBackup(backup.Path)) } +// writeLegacyDBFile creates a tmpo.db on disk with a pre-migration schema (no +// hourly_rate/milestone_name columns and no migration keys marked complete), +// simulating a database from before the migration system existed. +func writeLegacyDBFile(t *testing.T, path string) { + t.Helper() + db, err := sql.Open("sqlite", path) + assert.NoError(t, err) + defer db.Close() + + _, err = db.Exec(` + CREATE TABLE time_entries ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + project_name TEXT NOT NULL, + start_time DATETIME NOT NULL, + end_time DATETIME, + description TEXT + ) + `) + assert.NoError(t, err) + + _, err = db.Exec(` + CREATE TABLE milestones ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + project_name TEXT NOT NULL, + name TEXT NOT NULL, + start_time DATETIME NOT NULL, + end_time DATETIME, + UNIQUE(project_name, name) + ) + `) + assert.NoError(t, err) +} + +func TestInitialize_RemovesPreMigrationBackupOnSuccess(t *testing.T) { + tmpHome := t.TempDir() + t.Setenv("HOME", tmpHome) + t.Setenv("USERPROFILE", tmpHome) + t.Setenv("TMPO_DEV", "") + + // Seed a legacy database file with pending migrations + tmpoDir := filepath.Join(tmpHome, ".tmpo") + assert.NoError(t, os.MkdirAll(tmpoDir, 0700)) + writeLegacyDBFile(t, filepath.Join(tmpoDir, "tmpo.db")) + + db, err := Initialize() + assert.NoError(t, err) + defer db.Close() + + // Migrations succeeded, so the temporary pre-migration snapshot must be gone. + // The user should see no backups they did not create themselves. + backups, err := ListBackups() + assert.NoError(t, err) + assert.Empty(t, backups, "successful migration should leave no pre-migration backup behind") + + // Migrations should have completed and upgraded the schema + pending, err := db.hasPendingMigrations() + assert.NoError(t, err) + assert.False(t, pending) + assert.True(t, hasColumn(t, db, "time_entries", "milestone_name")) +} + +func TestInitialize_RetainsPreMigrationBackupOnFailure(t *testing.T) { + tmpHome := t.TempDir() + t.Setenv("HOME", tmpHome) + t.Setenv("USERPROFILE", tmpHome) + t.Setenv("TMPO_DEV", "") + + // Seed a legacy database whose start_time cannot be scanned as a timestamp, + // forcing the UTC migration to fail partway through. + tmpoDir := filepath.Join(tmpHome, ".tmpo") + assert.NoError(t, os.MkdirAll(tmpoDir, 0700)) + dbPath := filepath.Join(tmpoDir, "tmpo.db") + writeLegacyDBFile(t, dbPath) + + seed, err := sql.Open("sqlite", dbPath) + assert.NoError(t, err) + _, err = seed.Exec( + "INSERT INTO time_entries (project_name, start_time, description) VALUES (?, ?, ?)", + "broken-project", "not-a-timestamp", "corrupt row", + ) + assert.NoError(t, err) + assert.NoError(t, seed.Close()) + + // Initialize must fail, and the error must point at the retained snapshot + _, err = Initialize() + assert.Error(t, err) + assert.Contains(t, err.Error(), "pre-migration backup was preserved") + + // The snapshot must still be present for the user to restore from + backups, err := ListBackups() + assert.NoError(t, err) + assert.Len(t, backups, 1) + assert.True(t, strings.HasSuffix(backups[0].Filename, "-premigration.db")) +} + +func TestInitialize_NoBackupForFreshDB(t *testing.T) { + tmpHome := t.TempDir() + t.Setenv("HOME", tmpHome) + t.Setenv("USERPROFILE", tmpHome) + t.Setenv("TMPO_DEV", "") + + // No pre-existing database file: this is a brand-new install + db, err := Initialize() + assert.NoError(t, err) + defer db.Close() + + backups, err := ListBackups() + assert.NoError(t, err) + assert.Empty(t, backups, "fresh database should not trigger a pre-migration backup") +} + +func TestInitialize_NoBackupWhenAlreadyMigrated(t *testing.T) { + tmpHome := t.TempDir() + t.Setenv("HOME", tmpHome) + t.Setenv("USERPROFILE", tmpHome) + t.Setenv("TMPO_DEV", "") + + // First run creates and fully migrates the database + db, err := Initialize() + assert.NoError(t, err) + assert.NoError(t, db.Close()) + + // Second run against the up-to-date database must not create a backup + db2, err := Initialize() + assert.NoError(t, err) + defer db2.Close() + + backups, err := ListBackups() + assert.NoError(t, err) + assert.Empty(t, backups, "already-migrated database should not trigger a pre-migration backup") +} + func TestCreateBackup_SetsPrivatePermissions(t *testing.T) { if runtime.GOOS == "windows" { t.Skip("POSIX permissions are not enforced on Windows") diff --git a/internal/storage/db.go b/internal/storage/db.go index ffb105b..c57ee0f 100644 --- a/internal/storage/db.go +++ b/internal/storage/db.go @@ -3,8 +3,8 @@ package storage import ( "database/sql" "fmt" + "os" "path/filepath" - "strings" "time" "github.com/DylanDevelops/tmpo/internal/fsperm" @@ -31,12 +31,23 @@ func Initialize() (*Database, error) { } dbPath := filepath.Join(tmpoDir, "tmpo.db") + + _, statErr := os.Stat(dbPath) + dbExisted := statErr == nil + db, err := sql.Open("sqlite", dbPath) if err != nil { return nil, fmt.Errorf("failed to open database: %w", err) } + success := false + defer func() { + if !success { + db.Close() + } + }() + _, err = db.Exec(` CREATE TABLE IF NOT EXISTS time_entries ( id INTEGER PRIMARY KEY AUTOINCREMENT, @@ -44,7 +55,8 @@ func Initialize() (*Database, error) { start_time DATETIME NOT NULL, end_time DATETIME, description TEXT, - hourly_rate REAL + hourly_rate REAL, + milestone_name TEXT ) `) @@ -67,26 +79,6 @@ func Initialize() (*Database, error) { return nil, fmt.Errorf("failed to create milestones table: %w", err) } - _, err = db.Exec(`ALTER TABLE time_entries ADD COLUMN hourly_rate REAL`) - if err != nil && !isColumnExistsError(err) { - return nil, fmt.Errorf("failed to add hourly_rate column: %w", err) - } - - _, err = db.Exec(`ALTER TABLE time_entries ADD COLUMN milestone_name TEXT`) - if err != nil && !isColumnExistsError(err) { - return nil, fmt.Errorf("failed to add milestone_name column: %w", err) - } - - _, err = db.Exec(`CREATE INDEX IF NOT EXISTS idx_time_entries_milestone ON time_entries(milestone_name)`) - if err != nil { - return nil, fmt.Errorf("failed to create index: %w", err) - } - - _, err = db.Exec(`CREATE INDEX IF NOT EXISTS idx_milestones_project_active ON milestones(project_name, end_time)`) - if err != nil { - return nil, fmt.Errorf("failed to create index: %w", err) - } - // settings table for tracking migrations and other metadata _, err = db.Exec(` CREATE TABLE IF NOT EXISTS settings ( @@ -101,10 +93,34 @@ func Initialize() (*Database, error) { database := &Database{db: db} + var preMigrationBackupPath string + if dbExisted { + pending, err := database.hasPendingMigrations() + if err != nil { + return nil, fmt.Errorf("failed to check for pending migrations: %w", err) + } + if pending { + backup, err := database.createBackup("-premigration") + if err != nil { + return nil, fmt.Errorf("failed to create pre-migration backup: %w", err) + } + preMigrationBackupPath = backup.Path + } + } + if err := database.runMigrations(); err != nil { + if preMigrationBackupPath != "" { + return nil, fmt.Errorf("failed to run migrations (a pre-migration backup was preserved at %s): %w", preMigrationBackupPath, err) + } return nil, fmt.Errorf("failed to run migrations: %w", err) } + if preMigrationBackupPath != "" { + if err := os.Remove(preMigrationBackupPath); err != nil && !os.IsNotExist(err) { + return nil, fmt.Errorf("failed to remove temporary pre-migration backup: %w", err) + } + } + if err := fsperm.SecureFile(dbPath); err != nil { return nil, err } @@ -114,18 +130,10 @@ func Initialize() (*Database, error) { } } + success = true return database, nil } -func isColumnExistsError(err error) bool { - if err == nil { - return false - } - errMsg := err.Error() - return strings.Contains(errMsg, "duplicate column name") || - strings.Contains(errMsg, "duplicate column") -} - func (d *Database) CreateEntry(projectName, description string, hourlyRate *float64, milestoneName *string) (*TimeEntry, error) { var rate sql.NullFloat64 if hourlyRate != nil { diff --git a/internal/storage/migrations.go b/internal/storage/migrations.go index 21ef1e2..09d535c 100644 --- a/internal/storage/migrations.go +++ b/internal/storage/migrations.go @@ -10,11 +10,13 @@ import ( // ! I'm adding this system so that future database migrations will be easier - Dylan const ( Migration001_UTCTimestamps = "001_utc_timestamps" + Migration002_EntrySchema = "002_entry_schema" ) // allMigrationKeys lists every migration in order. Adding a new migration here automatically bumps CurrentSchemaVersion. var allMigrationKeys = []string{ Migration001_UTCTimestamps, + Migration002_EntrySchema, } // CurrentSchemaVersion is the number of migrations the current binary knows about. @@ -27,6 +29,136 @@ func (d *Database) runMigrations() error { return fmt.Errorf("timestamp UTC migration failed: %w", err) } + // Migration 2: Ensure time_entries has the hourly_rate/milestone_name + // columns and supporting indexes (previously ad-hoc ALTERs in Initialize). + if err := d.migrateEntrySchema(); err != nil { + return fmt.Errorf("entry schema migration failed: %w", err) + } + + return nil +} + +// hasPendingMigrations reports whether any known migration has not yet been +// marked complete. Used to decide whether a pre-migration backup is warranted. +func (d *Database) hasPendingMigrations() (bool, error) { + for _, key := range allMigrationKeys { + completed, err := d.hasMigrationRun(key) + if err != nil { + return false, err + } + if !completed { + return true, nil + } + } + return false, nil +} + +// columnExists reports whether the given column is present on the table by +// inspecting PRAGMA table_info. This replaces the old fragile error-string +// matching against "duplicate column name". +func columnExists(tx *sql.Tx, table, column string) (bool, error) { + rows, err := tx.Query(fmt.Sprintf("PRAGMA table_info(%s)", table)) + if err != nil { + return false, fmt.Errorf("failed to read table info for %s: %w", table, err) + } + defer rows.Close() + + for rows.Next() { + var ( + cid int + name string + ctype string + notnull int + dflt sql.NullString + pk int + ) + if err := rows.Scan(&cid, &name, &ctype, ¬null, &dflt, &pk); err != nil { + return false, fmt.Errorf("failed to scan table info for %s: %w", table, err) + } + if name == column { + return true, nil + } + } + + if err := rows.Err(); err != nil { + return false, fmt.Errorf("error iterating table info for %s: %w", table, err) + } + + return false, nil +} + +// migrateEntrySchema brings older databases up to the current time_entries +// schema by adding any missing columns and indexes. +func (d *Database) migrateEntrySchema() error { + completed, err := d.hasMigrationRun(Migration002_EntrySchema) + if err != nil { + return err + } + + if completed { + // migration is already finished + return nil + } + + // start transaction + tx, err := d.db.Begin() + if err != nil { + return fmt.Errorf("failed to begin transaction: %w", err) + } + + // rollback changes if something explodes + defer func() { + if err != nil { + tx.Rollback() + } + }() + + // add hourly_rate column if it is missing (legacy databases predate it) + var hasHourlyRate bool + if hasHourlyRate, err = columnExists(tx, "time_entries", "hourly_rate"); err != nil { + return err + } + if !hasHourlyRate { + if _, err = tx.Exec(`ALTER TABLE time_entries ADD COLUMN hourly_rate REAL`); err != nil { + return fmt.Errorf("failed to add hourly_rate column: %w", err) + } + } + + // add milestone_name column if it is missing + var hasMilestoneName bool + if hasMilestoneName, err = columnExists(tx, "time_entries", "milestone_name"); err != nil { + return err + } + if !hasMilestoneName { + if _, err = tx.Exec(`ALTER TABLE time_entries ADD COLUMN milestone_name TEXT`); err != nil { + return fmt.Errorf("failed to add milestone_name column: %w", err) + } + } + + // ensure supporting indexes exist + if _, err = tx.Exec(`CREATE INDEX IF NOT EXISTS idx_time_entries_milestone ON time_entries(milestone_name)`); err != nil { + return fmt.Errorf("failed to create milestone index: %w", err) + } + if _, err = tx.Exec(`CREATE INDEX IF NOT EXISTS idx_milestones_project_active ON milestones(project_name, end_time)`); err != nil { + return fmt.Errorf("failed to create milestones index: %w", err) + } + + // mark migration as complete in transaction + _, err = tx.Exec( + "INSERT OR REPLACE INTO settings (key, value, updated_at) VALUES (?, ?, ?)", + Migration002_EntrySchema, + "completed", + time.Now().UTC(), + ) + if err != nil { + return fmt.Errorf("failed to mark migration complete: %w", err) + } + + // push changes to db + if err = tx.Commit(); err != nil { + return fmt.Errorf("failed to commit migration transaction: %w", err) + } + return nil } diff --git a/internal/storage/migrations_test.go b/internal/storage/migrations_test.go index 05fff1b..a2fe082 100644 --- a/internal/storage/migrations_test.go +++ b/internal/storage/migrations_test.go @@ -427,6 +427,137 @@ func TestRunMigrations(t *testing.T) { assert.Equal(t, time.UTC, startTime.Location()) } +// setupLegacyEntriesTestDB creates an in-memory database whose time_entries +// table predates the hourly_rate and milestone_name columns, mimicking a very +// old tmpo installation. +func setupLegacyEntriesTestDB(t *testing.T) *Database { + db, err := sql.Open("sqlite", ":memory:") + assert.NoError(t, err) + + _, err = db.Exec(` + CREATE TABLE settings ( + key TEXT PRIMARY KEY, + value TEXT NOT NULL, + updated_at DATETIME NOT NULL + ) + `) + assert.NoError(t, err) + + _, err = db.Exec(` + CREATE TABLE time_entries ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + project_name TEXT NOT NULL, + start_time DATETIME NOT NULL, + end_time DATETIME, + description TEXT + ) + `) + assert.NoError(t, err) + + _, err = db.Exec(` + CREATE TABLE milestones ( + id INTEGER PRIMARY KEY AUTOINCREMENT, + project_name TEXT NOT NULL, + name TEXT NOT NULL, + start_time DATETIME NOT NULL, + end_time DATETIME, + UNIQUE(project_name, name) + ) + `) + assert.NoError(t, err) + + return &Database{db: db} +} + +// hasColumn is a test helper that reports whether a table has a given column. +func hasColumn(t *testing.T, db *Database, table, column string) bool { + t.Helper() + rows, err := db.db.Query("PRAGMA table_info(" + table + ")") + assert.NoError(t, err) + defer rows.Close() + + for rows.Next() { + var ( + cid int + name string + ctype string + notnull int + dflt sql.NullString + pk int + ) + assert.NoError(t, rows.Scan(&cid, &name, &ctype, ¬null, &dflt, &pk)) + if name == column { + return true + } + } + return false +} + +func TestMigrateEntrySchema_AddsMissingColumns(t *testing.T) { + db := setupLegacyEntriesTestDB(t) + defer db.Close() + + // Legacy schema lacks both columns + assert.False(t, hasColumn(t, db, "time_entries", "hourly_rate")) + assert.False(t, hasColumn(t, db, "time_entries", "milestone_name")) + + err := db.migrateEntrySchema() + assert.NoError(t, err) + + // Both columns should now exist + assert.True(t, hasColumn(t, db, "time_entries", "hourly_rate")) + assert.True(t, hasColumn(t, db, "time_entries", "milestone_name")) + + // Migration should be recorded as complete + hasRun, err := db.hasMigrationRun(Migration002_EntrySchema) + assert.NoError(t, err) + assert.True(t, hasRun) +} + +func TestMigrateEntrySchema_Idempotent(t *testing.T) { + db := setupLegacyEntriesTestDB(t) + defer db.Close() + + // First run adds the columns + assert.NoError(t, db.migrateEntrySchema()) + + // Second run must be a no-op and must not error on already-present columns + assert.NoError(t, db.migrateEntrySchema()) + + assert.True(t, hasColumn(t, db, "time_entries", "hourly_rate")) + assert.True(t, hasColumn(t, db, "time_entries", "milestone_name")) +} + +func TestMigrateEntrySchema_NoOpWhenColumnsPresent(t *testing.T) { + // setupMigrationTestDB already includes both columns (current schema) + db := setupMigrationTestDB(t) + defer db.Close() + + err := db.migrateEntrySchema() + assert.NoError(t, err) + + hasRun, err := db.hasMigrationRun(Migration002_EntrySchema) + assert.NoError(t, err) + assert.True(t, hasRun) +} + +func TestHasPendingMigrations(t *testing.T) { + db := setupMigrationTestDB(t) + defer db.Close() + + // Nothing marked complete yet + pending, err := db.hasPendingMigrations() + assert.NoError(t, err) + assert.True(t, pending) + + // Run every migration + assert.NoError(t, db.runMigrations()) + + pending, err = db.hasPendingMigrations() + assert.NoError(t, err) + assert.False(t, pending) +} + func TestMigrateTimeEntriesTableToUTC_EmptyTable(t *testing.T) { db := setupMigrationTestDB(t) defer db.Close()