diff --git a/replica.go b/replica.go index 2584abd34..34422e426 100644 --- a/replica.go +++ b/replica.go @@ -738,12 +738,11 @@ func (r *Replica) Restore(ctx context.Context, opt RestoreOptions) (err error) { pr, pw := io.Pipe() go func() { - c, err := ltx.NewCompactor(pw, rdrs) + c, err := newRestoreCompactor(pw, rdrs) if err != nil { pw.CloseWithError(fmt.Errorf("new ltx compactor: %w", err)) return } - c.HeaderFlags = ltx.HeaderFlagNoChecksum _ = pw.CloseWithError(c.Compact(ctx)) }() @@ -798,6 +797,32 @@ func (r *Replica) Restore(ctx context.Context, opt RestoreOptions) (err error) { return nil } +func newRestoreCompactor(w io.Writer, rdrs []io.Reader) (*ltx.Compactor, error) { + if len(rdrs) == 0 { + return nil, fmt.Errorf("at least one input reader required") + } + + last := len(rdrs) - 1 + finalHeader, replay, err := ltx.PeekHeader(rdrs[last]) + if err != nil { + return nil, fmt.Errorf("peek final input header: %w", err) + } + if closer, ok := rdrs[last].(io.Closer); ok { + rdrs[last] = internal.NewReadCloser(replay, closer) + } else { + rdrs[last] = replay + } + + c, err := ltx.NewCompactor(w, rdrs) + if err != nil { + return nil, err + } + if finalHeader.NoChecksum() { + c.HeaderFlags = ltx.HeaderFlagNoChecksum + } + return c, nil +} + // follow enters a continuous restore loop, polling for new LTX files and // applying them to the restored database. It blocks until the context is // cancelled (e.g. Ctrl+C). Returns nil on clean shutdown. diff --git a/replica_test.go b/replica_test.go index cfa9677d8..1aebba283 100644 --- a/replica_test.go +++ b/replica_test.go @@ -276,6 +276,85 @@ func TestReplica_RestoreRetriesInitialLTXOpenError(t *testing.T) { } } +func TestReplica_Restore_ChecksumTracked(t *testing.T) { + const pageSize = 512 + + ctx := context.Background() + client := file.NewReplicaClient(t.TempDir()) + + page1 := bytes.Repeat([]byte{0x11}, pageSize) + page2 := bytes.Repeat([]byte{0x22}, pageSize) + updatedPage2 := bytes.Repeat([]byte{0x33}, pageSize) + + snapshotChecksum := ltx.ChecksumFlag + snapshotChecksum = ltx.ChecksumFlag | (snapshotChecksum ^ ltx.ChecksumPage(1, page1)) + snapshotChecksum = ltx.ChecksumFlag | (snapshotChecksum ^ ltx.ChecksumPage(2, page2)) + writeChecksumTrackedLTXFile(t, client, litestream.SnapshotLevel, ltx.Header{ + Version: ltx.Version, + PageSize: pageSize, + Commit: 2, + MinTXID: 1, + MaxTXID: 1, + Timestamp: 1000, + }, []uint32{1, 2}, [][]byte{page1, page2}, snapshotChecksum) + + updatedChecksum := ltx.ChecksumFlag | (snapshotChecksum ^ ltx.ChecksumPage(2, page2) ^ ltx.ChecksumPage(2, updatedPage2)) + writeChecksumTrackedLTXFile(t, client, 0, ltx.Header{ + Version: ltx.Version, + PageSize: pageSize, + Commit: 2, + MinTXID: 2, + MaxTXID: 2, + Timestamp: 2000, + PreApplyChecksum: snapshotChecksum, + }, []uint32{2}, [][]byte{updatedPage2}, updatedChecksum) + + r := litestream.NewReplicaWithClient(nil, client) + restorePath := filepath.Join(t.TempDir(), "restored.db") + if err := r.Restore(ctx, litestream.RestoreOptions{OutputPath: restorePath}); err != nil { + t.Fatal(err) + } + + got, err := os.ReadFile(restorePath) + if err != nil { + t.Fatal(err) + } + want := append(append([]byte(nil), page1...), updatedPage2...) + if !bytes.Equal(got, want) { + t.Fatalf("restored database mismatch: got %d bytes, want %d bytes", len(got), len(want)) + } +} + +func writeChecksumTrackedLTXFile(t testing.TB, client litestream.ReplicaClient, level int, hdr ltx.Header, pgnos []uint32, pages [][]byte, postApplyChecksum ltx.Checksum) { + t.Helper() + + var buf bytes.Buffer + enc, err := ltx.NewEncoder(&buf) + if err != nil { + t.Fatal(err) + } + if err := enc.EncodeHeader(hdr); err != nil { + t.Fatal(err) + } + for i, pgno := range pgnos { + if err := enc.EncodePage(ltx.PageHeader{Pgno: pgno}, pages[i]); err != nil { + t.Fatal(err) + } + } + enc.SetPostApplyChecksum(postApplyChecksum) + if err := enc.Close(); err != nil { + t.Fatal(err) + } + + data := append([]byte(nil), buf.Bytes()...) + if err := ltx.NewDecoder(bytes.NewReader(data)).Verify(); err != nil { + t.Fatal(err) + } + if _, err := client.WriteLTXFile(context.Background(), level, hdr.MinTXID, hdr.MaxTXID, bytes.NewReader(data)); err != nil { + t.Fatal(err) + } +} + type transientOpenFailureClient struct { litestream.ReplicaClient remainingFailures int diff --git a/vfs.go b/vfs.go index bf2a96d65..40e05dedc 100644 --- a/vfs.go +++ b/vfs.go @@ -699,11 +699,10 @@ func (h *Hydrator) Restore(ctx context.Context, infos []*ltx.FileInfo) error { // Compact and decode using io.Pipe pattern pr, pw := io.Pipe() - c, err := ltx.NewCompactor(pw, rdrs) + c, err := newRestoreCompactor(pw, rdrs) if err != nil { return fmt.Errorf("new ltx compactor: %w", err) } - c.HeaderFlags = ltx.HeaderFlagNoChecksum h.compactor = c go func() { diff --git a/vfs_test.go b/vfs_test.go index 23db0bad6..d59a56154 100644 --- a/vfs_test.go +++ b/vfs_test.go @@ -849,6 +849,8 @@ func newCountingReplicaClient() *countingReplicaClient { return &countingReplica func (c *countingReplicaClient) Type() string { return "count" } +func (c *countingReplicaClient) SetLogger(*slog.Logger) {} + func (c *countingReplicaClient) Init(context.Context) error { return nil } func (c *countingReplicaClient) LTXFiles(ctx context.Context, level int, seek ltx.TXID, useMetadata bool) (ltx.FileIterator, error) { @@ -881,6 +883,8 @@ func newBlockingReplicaClient() *blockingReplicaClient { func (c *mockReplicaClient) Type() string { return "mock" } +func (c *mockReplicaClient) SetLogger(*slog.Logger) {} + func (c *mockReplicaClient) Init(context.Context) error { return nil } func (c *mockReplicaClient) addFixture(tb testing.TB, fx *ltxFixture) { @@ -1045,6 +1049,100 @@ func buildLTXFixtureWithPages(tb testing.TB, txid ltx.TXID, pageSize uint32, pgn return <xFixture{info: info, data: buf.Bytes()} } +func buildChecksumTrackedLTXFixture(tb testing.TB, hdr ltx.Header, pgnos []uint32, pages [][]byte, postApplyChecksum ltx.Checksum) *ltxFixture { + tb.Helper() + + var buf bytes.Buffer + enc, err := ltx.NewEncoder(&buf) + if err != nil { + tb.Fatalf("new encoder: %v", err) + } + if err := enc.EncodeHeader(hdr); err != nil { + tb.Fatalf("encode header: %v", err) + } + for i, pgno := range pgnos { + if err := enc.EncodePage(ltx.PageHeader{Pgno: pgno}, pages[i]); err != nil { + tb.Fatalf("encode page %d: %v", pgno, err) + } + } + enc.SetPostApplyChecksum(postApplyChecksum) + if err := enc.Close(); err != nil { + tb.Fatalf("close encoder: %v", err) + } + + data := append([]byte(nil), buf.Bytes()...) + if err := ltx.NewDecoder(bytes.NewReader(data)).Verify(); err != nil { + tb.Fatalf("verify fixture: %v", err) + } + return <xFixture{ + info: <x.FileInfo{ + Level: 0, + MinTXID: hdr.MinTXID, + MaxTXID: hdr.MaxTXID, + Size: int64(len(data)), + CreatedAt: time.UnixMilli(hdr.Timestamp).UTC(), + }, + data: data, + } +} + +func TestHydrator_Restore_ChecksumTracked(t *testing.T) { + const pageSize = 512 + + client := newMockReplicaClient() + page1 := bytes.Repeat([]byte{0x11}, pageSize) + page2 := bytes.Repeat([]byte{0x22}, pageSize) + updatedPage2 := bytes.Repeat([]byte{0x33}, pageSize) + + snapshotChecksum := ltx.ChecksumFlag + snapshotChecksum = ltx.ChecksumFlag | (snapshotChecksum ^ ltx.ChecksumPage(1, page1)) + snapshotChecksum = ltx.ChecksumFlag | (snapshotChecksum ^ ltx.ChecksumPage(2, page2)) + snapshot := buildChecksumTrackedLTXFixture(t, ltx.Header{ + Version: ltx.Version, + PageSize: pageSize, + Commit: 2, + MinTXID: 1, + MaxTXID: 1, + Timestamp: 1000, + }, []uint32{1, 2}, [][]byte{page1, page2}, snapshotChecksum) + snapshot.info.Level = SnapshotLevel + client.addFixture(t, snapshot) + + updatedChecksum := ltx.ChecksumFlag | (snapshotChecksum ^ ltx.ChecksumPage(2, page2) ^ ltx.ChecksumPage(2, updatedPage2)) + incremental := buildChecksumTrackedLTXFixture(t, ltx.Header{ + Version: ltx.Version, + PageSize: pageSize, + Commit: 2, + MinTXID: 2, + MaxTXID: 2, + Timestamp: 2000, + PreApplyChecksum: snapshotChecksum, + }, []uint32{2}, [][]byte{updatedPage2}, updatedChecksum) + client.addFixture(t, incremental) + + h := NewHydrator(filepath.Join(t.TempDir(), "hydration.db"), false, pageSize, client, slog.Default()) + if err := h.Init(); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = h.Close() }) + + if err := h.Restore(context.Background(), []*ltx.FileInfo{snapshot.info, incremental.info}); err != nil { + t.Fatal(err) + } + if got, want := h.TXID(), incremental.info.MaxTXID; got != want { + t.Fatalf("txid=%s, want %s", got, want) + } + + got := make([]byte, 2*pageSize) + if _, err := h.file.ReadAt(got, 0); err != nil { + t.Fatal(err) + } + want := append(append([]byte(nil), page1...), updatedPage2...) + if !bytes.Equal(got, want) { + t.Fatalf("hydrated database mismatch: got %d bytes, want %d bytes", len(got), len(want)) + } +} + // TestVFSFile_Hydration_Basic tests that hydration completes and reads from local file. func TestVFSFile_Hydration_Basic(t *testing.T) { client := newMockReplicaClient() diff --git a/vfs_write_test.go b/vfs_write_test.go index 6afb5dcef..06243c880 100644 --- a/vfs_write_test.go +++ b/vfs_write_test.go @@ -37,6 +37,8 @@ func newWriteTestReplicaClient() *writeTestReplicaClient { func (c *writeTestReplicaClient) Type() string { return "test" } +func (c *writeTestReplicaClient) SetLogger(*slog.Logger) {} + func (c *writeTestReplicaClient) Init(ctx context.Context) error { return nil } func (c *writeTestReplicaClient) LTXFiles(ctx context.Context, level int, seek ltx.TXID, useMetadata bool) (ltx.FileIterator, error) {