| 1 | package draftstate |
| 2 | |
| 3 | import ( |
| 4 | "bytes" |
| 5 | "database/sql" |
| 6 | "errors" |
| 7 | "fmt" |
| 8 | "os" |
| 9 | "path/filepath" |
| 10 | "testing" |
| 11 | ) |
| 12 | |
| 13 | func TestUnrecognizedDraftDatabaseIsNeverInitializedOrRewritten(t *testing.T) { |
| 14 | for _, version := range []int{-1, 0, SchemaVersion + 1} { |
| 15 | t.Run(fmt.Sprint(version), func(t *testing.T) { |
| 16 | path := filepath.Join(t.TempDir(), "drafts.sqlite") |
| 17 | db, err := sql.Open("sqlite", path) |
| 18 | if err != nil { |
| 19 | t.Fatal(err) |
| 20 | } |
| 21 | for _, query := range []string{`CREATE TABLE future_input(body TEXT)`, `INSERT INTO future_input VALUES('unsent content')`, fmt.Sprintf("PRAGMA user_version=%d", version)} { |
| 22 | if _, err := db.Exec(query); err != nil { |
| 23 | t.Fatal(err) |
| 24 | } |
| 25 | } |
| 26 | if err := db.Close(); err != nil { |
| 27 | t.Fatal(err) |
| 28 | } |
| 29 | before, err := os.ReadFile(path) |
| 30 | if err != nil { |
| 31 | t.Fatal(err) |
| 32 | } |
| 33 | for range 2 { |
| 34 | store := New(path) |
| 35 | _, err := store.ListActive(t.Context()) |
| 36 | _ = store.Close() |
| 37 | if !errors.Is(err, ErrUnsupportedVersion) { |
| 38 | t.Fatalf("unrecognized database should be blocked, got %v", err) |
| 39 | } |
| 40 | after, err := os.ReadFile(path) |
| 41 | if err != nil || !bytes.Equal(before, after) { |
| 42 | t.Fatalf("unrecognized database was rewritten: %v", err) |
| 43 | } |
| 44 | } |
| 45 | }) |
| 46 | } |
| 47 | } |
| 48 |