Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion internal/storage/backup.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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, "'", "''")
Expand Down
132 changes: 132 additions & 0 deletions internal/storage/backup_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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")
Expand Down
70 changes: 39 additions & 31 deletions internal/storage/db.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,8 @@ package storage
import (
"database/sql"
"fmt"
"os"
"path/filepath"
"strings"
"time"

"github.com/DylanDevelops/tmpo/internal/fsperm"
Expand All @@ -31,20 +31,32 @@ 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,
project_name TEXT NOT NULL,
start_time DATETIME NOT NULL,
end_time DATETIME,
description TEXT,
hourly_rate REAL
hourly_rate REAL,
milestone_name TEXT
)
`)

Expand All @@ -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 (
Expand All @@ -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
}
Expand All @@ -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 {
Expand Down
Loading