From 349063968b262c1710b39d2287a1f3161814488f Mon Sep 17 00:00:00 2001 From: dailz Date: Fri, 12 Jun 2026 13:43:24 +0800 Subject: [PATCH] feat(wal,memtable): implement segment writer, block writer, and MemTable - wal/block_writer.go: 32KB block buffer with padding and flush - wal/segment_writer.go: WAL segment file with durable-ready protocol - memtable/memtable.go: Arena+SkipList wrapper with publish/abort semantics - Comprehensive tests for all modules, all pass with -race --- .omo/boulder.json | 30 +++- .omo/plans/phase1-wal.md | 4 +- memtable/memtable.go | 159 +++++++++++++++++ memtable/memtable_test.go | 209 ++++++++++++++++++++++ wal/block_writer.go | 147 +++++++++++++++ wal/block_writer_test.go | 353 +++++++++++++++++++++++++++++++++++++ wal/segment_writer.go | 170 ++++++++++++++++++ wal/segment_writer_test.go | 339 +++++++++++++++++++++++++++++++++++ 8 files changed, 1407 insertions(+), 4 deletions(-) create mode 100644 memtable/memtable.go create mode 100644 memtable/memtable_test.go create mode 100644 wal/block_writer.go create mode 100644 wal/block_writer_test.go create mode 100644 wal/segment_writer.go create mode 100644 wal/segment_writer_test.go diff --git a/.omo/boulder.json b/.omo/boulder.json index c8c38ac..216ef43 100644 --- a/.omo/boulder.json +++ b/.omo/boulder.json @@ -8,7 +8,7 @@ "plan_name": "phase1-wal", "status": "active", "started_at": "2026-06-12T05:09:32.588Z", - "updated_at": "2026-06-12T05:33:52.854Z", + "updated_at": "2026-06-12T05:43:11.052Z", "session_ids": [ "opencode:ses_145c3bae9ffeTB2zbsTym0Cev8" ], @@ -68,6 +68,19 @@ "status": "completed", "ended_at": "2026-06-12T05:33:52.854Z", "elapsed_ms": 64624 + }, + "todo:10": { + "task_key": "todo:10", + "task_label": "10", + "task_title": "Segment writer (file write + block writer)", + "session_id": "opencode:ses_145acc017ffe2fra67VT2L28LQ", + "agent": "Sisyphus-Junior", + "category": "unspecified-high", + "updated_at": "2026-06-12T05:43:11.052Z", + "started_at": "2026-06-12T05:42:38.284Z", + "status": "completed", + "ended_at": "2026-06-12T05:43:11.052Z", + "elapsed_ms": 32768 } } } @@ -75,7 +88,7 @@ "active_plan": "/home/dailz/workspace/src/go-kv/.omo/plans/phase1-wal.md", "started_at": "2026-06-12T05:09:32.588Z", "status": "active", - "updated_at": "2026-06-12T05:33:52.854Z", + "updated_at": "2026-06-12T05:43:11.052Z", "session_ids": [ "opencode:ses_145c3bae9ffeTB2zbsTym0Cev8" ], @@ -135,6 +148,19 @@ "status": "completed", "ended_at": "2026-06-12T05:33:52.854Z", "elapsed_ms": 64624 + }, + "todo:10": { + "task_key": "todo:10", + "task_label": "10", + "task_title": "Segment writer (file write + block writer)", + "session_id": "opencode:ses_145acc017ffe2fra67VT2L28LQ", + "agent": "Sisyphus-Junior", + "category": "unspecified-high", + "updated_at": "2026-06-12T05:43:11.052Z", + "started_at": "2026-06-12T05:42:38.284Z", + "status": "completed", + "ended_at": "2026-06-12T05:43:11.052Z", + "elapsed_ms": 32768 } }, "agent": "atlas" diff --git a/.omo/plans/phase1-wal.md b/.omo/plans/phase1-wal.md index ed86a43..5002700 100644 --- a/.omo/plans/phase1-wal.md +++ b/.omo/plans/phase1-wal.md @@ -823,7 +823,7 @@ Max Concurrent: 6 (Wave 1b) **Commit**: YES (group with Task 6) -- [ ] 10. Segment writer (file write + block writer) +- [x] 10. Segment writer (file write + block writer) **What to do**: - 创建 `wal/block_writer.go`:`BlockWriter` — 管理 32KB block 的 Physical Record 写入 @@ -1105,7 +1105,7 @@ Max Concurrent: 6 (Wave 1b) **Commit**: YES (with Task 12) - Message: `feat(wal): implement WAL writer with group commit` -- [ ] 14. MemTable (wrap skiplist + arena + publish/abort) +- [x] 14. MemTable (wrap skiplist + arena + publish/abort) **What to do**: - 创建 `memtable/memtable.go`:`MemTable` 结构体 diff --git a/memtable/memtable.go b/memtable/memtable.go new file mode 100644 index 0000000..d96f1a0 --- /dev/null +++ b/memtable/memtable.go @@ -0,0 +1,159 @@ +package memtable + +import ( + "errors" + "sync" + "sync/atomic" +) + +// ErrMemTableFull is returned when the memtable cannot reserve enough space. +var ErrMemTableFull = errors.New("memtable: insufficient capacity") + +// ReserveEntry describes a single key-value pair to be reserved. +type ReserveEntry struct { + Key []byte + Value []byte + IsDelete bool +} + +// GetResult is returned by MemTable.Get. +type GetResult struct { + Found bool + Value []byte + Sequence uint64 +} + +// MemTable wraps an Arena and SkipList with publish/abort semantics. +// +// Writers first reserve capacity, then insert pending entries, then either +// publish (making them visible) or abort (leaving them invisible forever). +// Readers only see published entries that have not been aborted. +type MemTable struct { + arena *Arena + skiplist *SkipList + published atomic.Uint64 // highest published sequence + aborted sync.Map // map[uint64]struct{} — set of aborted sequences +} + +// NewMemTable creates a MemTable backed by an arena of the given capacity. +func NewMemTable(capacity uint32) *MemTable { + arena := NewArena(capacity) + return &MemTable{ + arena: arena, + skiplist: NewSkipList(arena), + } +} + +// Reserve estimates whether the arena has enough space for all entries. +// It returns the total estimated size for caller tracking. +// +// The estimate is conservative: key + value bytes plus a fixed metadata +// overhead per entry. Since Phase 1 stores nodes on the Go heap (not in +// the arena), this reservation is primarily a capacity gate for future +// arena-backed phases. +func (mt *MemTable) Reserve(entries []ReserveEntry) (uint32, error) { + const metadataOverhead uint32 = 32 // per-entry overhead estimate + + var totalSize uint32 + for _, e := range entries { + sz := metadataOverhead + uint32(len(e.Key)) + uint32(len(e.Value)) + // Align to 8 bytes + sz = (sz + 7) &^ uint32(7) + totalSize += sz + } + + if err := mt.arena.Reserve(totalSize); err != nil { + return 0, ErrMemTableFull + } + return totalSize, nil +} + +// PutPending inserts a key-value pair as pending (invisible to Get/Iterator). +func (mt *MemTable) PutPending(key []byte, value []byte, sequence uint64) error { + return mt.skiplist.Put(key, value, sequence, true) +} + +// DeletePending inserts a tombstone entry (nil value) as pending. +func (mt *MemTable) DeletePending(key []byte, sequence uint64) error { + return mt.skiplist.Put(key, nil, sequence, true) +} + +// Publish makes all pending entries with sequence <= upToSequence visible. +// +// It iterates the skip list and re-inserts each pending entry with +// pending=false, which the lock-free reader path will then observe. +// Aborted entries are skipped and remain invisible forever. +func (mt *MemTable) Publish(upToSequence uint64) { + // Collect entries to publish under the iterator (which skips pending). + // We need a raw walk, so we access the skiplist directly. + type entry struct { + key []byte + value []byte + sequence uint64 + } + + // Walk the raw skip list level-0 chain to find pending entries. + // We cannot use NewIterator because it skips pending entries. + var toPublish []entry + + mt.skiplist.mu.Lock() + node := mt.skiplist.head.next[0].Load() + for node != nil { + if node.pending && node.sequence <= upToSequence { + if _, aborted := mt.aborted.Load(node.sequence); !aborted { + toPublish = append(toPublish, entry{ + key: node.key, + value: node.value, + sequence: node.sequence, + }) + } + } + node = node.next[0].Load() + } + mt.skiplist.mu.Unlock() + + // Re-put each entry as published. Each call acquires the skiplist mutex. + for _, e := range toPublish { + _ = mt.skiplist.Put(e.key, e.value, e.sequence, false) + } + + // Update high-water mark after all entries are visible. + mt.published.Store(upToSequence) +} + +// Abort marks a sequence as aborted. Aborted entries are never visible to readers. +func (mt *MemTable) Abort(sequence uint64) { + mt.aborted.Store(sequence, struct{}{}) +} + +// Get retrieves the latest visible value for key. +// +// A value is visible if it is published (pending=false in the skip list) +// and not in the aborted set. +func (mt *MemTable) Get(key []byte) *GetResult { + found, value, sequence := mt.skiplist.Get(key) + if !found { + return &GetResult{Found: false} + } + return &GetResult{ + Found: true, + Value: value, + Sequence: sequence, + } +} + +// NewIterator returns a forward iterator over all published (visible) entries. +// Pending and aborted entries are automatically skipped. +func (mt *MemTable) NewIterator() *Iterator { + return mt.skiplist.NewIterator() +} + +// ApproximateSize returns the number of bytes used in the arena. +func (mt *MemTable) ApproximateSize() uint64 { + return uint64(mt.arena.Capacity() - mt.arena.Remaining()) +} + +// UsableCapacity returns the remaining bytes available in the arena. +func (mt *MemTable) UsableCapacity() uint32 { + return mt.arena.Remaining() +} diff --git a/memtable/memtable_test.go b/memtable/memtable_test.go new file mode 100644 index 0000000..b059c38 --- /dev/null +++ b/memtable/memtable_test.go @@ -0,0 +1,209 @@ +package memtable + +import ( + "sync" + "testing" +) + +func TestMemTablePendingInvisible(t *testing.T) { + mt := NewMemTable(4096) + + // Put a pending entry — must not be visible. + if err := mt.PutPending([]byte("key1"), []byte("val1"), 1); err != nil { + t.Fatalf("PutPending: %v", err) + } + + res := mt.Get([]byte("key1")) + if res.Found { + t.Fatal("pending entry should not be visible") + } + + // Publish sequence 1 — entry becomes visible. + mt.Publish(1) + + res = mt.Get([]byte("key1")) + if !res.Found { + t.Fatal("published entry should be visible") + } + if string(res.Value) != "val1" { + t.Fatalf("value mismatch: got %q, want %q", res.Value, "val1") + } + if res.Sequence != 1 { + t.Fatalf("sequence mismatch: got %d, want %d", res.Sequence, 1) + } +} + +func TestMemTableAbortedInvisible(t *testing.T) { + mt := NewMemTable(4096) + + if err := mt.PutPending([]byte("key1"), []byte("val1"), 1); err != nil { + t.Fatalf("PutPending: %v", err) + } + + // Abort sequence 1 before publishing. + mt.Abort(1) + + // Publish up to sequence 10 — aborted entry should stay invisible. + mt.Publish(10) + + res := mt.Get([]byte("key1")) + if res.Found { + t.Fatal("aborted entry should not be visible") + } +} + +func TestMemTableMultiplePublish(t *testing.T) { + mt := NewMemTable(8192) + + if err := mt.PutPending([]byte("a"), []byte("va"), 1); err != nil { + t.Fatalf("PutPending a: %v", err) + } + if err := mt.PutPending([]byte("b"), []byte("vb"), 2); err != nil { + t.Fatalf("PutPending b: %v", err) + } + + // Neither visible before publish. + if mt.Get([]byte("a")).Found { + t.Fatal("a should not be visible before publish") + } + if mt.Get([]byte("b")).Found { + t.Fatal("b should not be visible before publish") + } + + // Publish both. + mt.Publish(2) + + if !mt.Get([]byte("a")).Found { + t.Fatal("a should be visible after publish") + } + if !mt.Get([]byte("b")).Found { + t.Fatal("b should be visible after publish") + } + + if string(mt.Get([]byte("a")).Value) != "va" { + t.Fatal("value mismatch for a") + } + if string(mt.Get([]byte("b")).Value) != "vb" { + t.Fatal("value mismatch for b") + } +} + +func TestMemTableReserve(t *testing.T) { + mt := NewMemTable(4096) + + // Reserve 10 small entries — should succeed. + entries := make([]ReserveEntry, 10) + for i := range entries { + entries[i] = ReserveEntry{ + Key: []byte("k"), + Value: []byte("v"), + } + } + total, err := mt.Reserve(entries) + if err != nil { + t.Fatalf("Reserve 10 entries: %v", err) + } + if total == 0 { + t.Fatal("expected non-zero total size") + } + + // Reserve more than remaining — should fail. + big := []ReserveEntry{{ + Key: make([]byte, 2048), + Value: make([]byte, 2048), + }} + _, err = mt.Reserve(big) + if err == nil { + t.Fatal("expected ErrMemTableFull for oversized reserve") + } +} + +func TestMemTableConcurrentReads(t *testing.T) { + mt := NewMemTable(8192) + + // Write 10 pending entries. + for i := 0; i < 10; i++ { + key := []byte{byte('a' + i)} + val := []byte{byte(i)} + if err := mt.PutPending(key, val, uint64(i+1)); err != nil { + t.Fatalf("PutPending %d: %v", i, err) + } + } + + // Publish all. + mt.Publish(10) + + // Concurrent reads should all see published values. + var wg sync.WaitGroup + for g := 0; g < 4; g++ { + wg.Add(1) + go func() { + defer wg.Done() + for i := 0; i < 100; i++ { + key := []byte{byte('a' + (i % 10))} + res := mt.Get(key) + if !res.Found { + t.Errorf("key %s not found", key) + return + } + } + }() + } + wg.Wait() +} + +func TestMemTableIteratorSkipsPending(t *testing.T) { + mt := NewMemTable(4096) + + _ = mt.PutPending([]byte("pending"), []byte("p"), 1) + _ = mt.PutPending([]byte("published"), []byte("pub"), 2) + mt.Publish(2) // publishes "published" at seq 2 + + // Abort seq 1 — stays pending. + // Actually we already published up to 2 which would publish seq 1 too. + // Let's test differently: insert a pending entry after publish. + _ = mt.PutPending([]byte("still_pending"), []byte("sp"), 3) + + it := mt.NewIterator() + count := 0 + for it.Valid() { + count++ + it.Next() + } + + // Should see "pending" (seq 1, now published) and "published" (seq 2), + // but not "still_pending" (seq 3, still pending). + if count != 2 { + t.Fatalf("expected 2 visible entries, got %d", count) + } +} + +func TestMemTableDeletePending(t *testing.T) { + mt := NewMemTable(4096) + + // Insert then delete. + _ = mt.PutPending([]byte("key1"), []byte("val1"), 1) + _ = mt.DeletePending([]byte("key1"), 2) + + mt.Publish(2) + + // After publish, the latest entry is a tombstone (nil value). + res := mt.Get([]byte("key1")) + if !res.Found { + t.Fatal("tombstone entry should still be 'found'") + } + if res.Value != nil { + t.Fatalf("expected nil value for tombstone, got %q", res.Value) + } +} + +func TestMemTableApproximateSize(t *testing.T) { + mt := NewMemTable(4096) + + if mt.ApproximateSize() != 0 { + t.Fatal("new memtable should have zero approximate size") + } + if mt.UsableCapacity() != 4096 { + t.Fatalf("usable capacity: got %d, want 4096", mt.UsableCapacity()) + } +} diff --git a/wal/block_writer.go b/wal/block_writer.go new file mode 100644 index 0000000..3d5b606 --- /dev/null +++ b/wal/block_writer.go @@ -0,0 +1,147 @@ +package wal + +import ( + "errors" + "fmt" + "io" +) + +// BlockWriter manages a single 32 KB block buffer for writing physical records. +// It handles block boundary padding and flushing complete blocks to an io.Writer. +type BlockWriter struct { + buf [WalBlockSize]byte + offset uint32 // current write position within the block +} + +// NewBlockWriter creates a BlockWriter ready to write into a fresh block. +func NewBlockWriter() *BlockWriter { + return &BlockWriter{} +} + +// BlockOffset returns the current write offset within the block (0..WalBlockSize). +func (bw *BlockWriter) BlockOffset() uint32 { + return bw.offset +} + +// WriteRecord writes a single physical record into the block buffer. +// If the record (header + payload) does not fit in the remaining space, +// the current block is padded with zeros and flushed to w, then the record +// is written at the start of a fresh block. +// +// Precondition: payload length must be ≤ WalBlockSize - PhysicalRecordHeaderSize +// (the caller is responsible for splitting large batches into appropriately-sized chunks). +func (bw *BlockWriter) WriteRecord(recType uint8, payload []byte, w io.Writer) error { + recordSize := PhysicalRecordHeaderSize + len(payload) + + if recordSize > WalBlockSize { + return fmt.Errorf("wal: record size %d exceeds block size %d", + recordSize, WalBlockSize) + } + + // Check if padding is needed before writing this record. + pad := bw.paddingNeeded() + if pad > 0 { + // Pad remaining bytes with zeros and flush. + if err := bw.flushPadded(w, pad); err != nil { + return fmt.Errorf("wal: flushing padded block: %w", err) + } + } + + // Check if the record fits in the current block. + remaining := WalBlockSize - bw.offset + if uint32(recordSize) > remaining { + // Not enough room — pad the rest and flush, then start a new block. + pad = int(remaining) + if err := bw.flushPadded(w, pad); err != nil { + return fmt.Errorf("wal: flushing partial block: %w", err) + } + } + + // Encode physical record directly into the block buffer. + encoded := EncodePhysicalRecord(recType, payload) + copy(bw.buf[bw.offset:], encoded) + bw.offset += uint32(len(encoded)) + + // If the block is exactly full, flush it immediately. + if bw.offset == WalBlockSize { + if _, err := w.Write(bw.buf[:]); err != nil { + return fmt.Errorf("wal: writing full block: %w", err) + } + bw.offset = 0 + } + + return nil +} + +// Flush writes the current block buffer to w, padding unused bytes with zeros. +// If the block is empty (offset == 0), this is a no-op. +func (bw *BlockWriter) Flush(w io.Writer) error { + if bw.offset == 0 { + return nil + } + return bw.flushPadded(w, int(WalBlockSize-bw.offset)) +} + +// Reset clears the block buffer, returning it to an empty state. +func (bw *BlockWriter) Reset() { + bw.offset = 0 + // Zero the buffer so partial blocks are padded with zeros. + for i := range bw.buf { + bw.buf[i] = 0 + } +} + +// paddingNeeded returns the number of zero-padding bytes required at the current +// block offset. When the remaining space in the block is ≤ PhysicalRecordHeaderSize (7), +// that space cannot hold even a minimal physical record and must be zero-padded. +func (bw *BlockWriter) paddingNeeded() int { + remaining := WalBlockSize - bw.offset + if remaining <= PhysicalRecordHeaderSize { + return int(remaining) + } + return 0 +} + +// flushPadded pads the remaining bytes with zeros and writes the full block to w. +// pad is the number of trailing bytes to zero-fill (WalBlockSize - offset - pad already zero +// from initial state or previous Reset). +func (bw *BlockWriter) flushPadded(w io.Writer, pad int) error { + if pad <= 0 { + return nil + } + + // Zero-fill padding region. The buffer was zeroed at init/reset, + // but we write explicitly for safety after partial record writes. + for i := uint32(0); i < uint32(pad); i++ { + bw.buf[bw.offset+i] = 0 + } + + if _, err := w.Write(bw.buf[:]); err != nil { + return err + } + + bw.offset = 0 + for i := range bw.buf { + bw.buf[i] = 0 + } + + return nil +} + +// Bytes returns a copy of the current block contents up to the current offset. +// Useful for testing. +func (bw *BlockWriter) Bytes() []byte { + out := make([]byte, bw.offset) + copy(out, bw.buf[:bw.offset]) + return out +} + +// FullBlockBytes returns the full block buffer. Only valid when offset == WalBlockSize. +func (bw *BlockWriter) FullBlockBytes() []byte { + out := make([]byte, WalBlockSize) + copy(out, bw.buf[:]) + return out +} + +// errBlockWriterNil is returned when a nil writer is passed to write operations. +var errBlockWriterNil = errors.New("wal: writer must not be nil") diff --git a/wal/block_writer_test.go b/wal/block_writer_test.go new file mode 100644 index 0000000..38ad3c8 --- /dev/null +++ b/wal/block_writer_test.go @@ -0,0 +1,353 @@ +package wal + +import ( + "bytes" + "encoding/binary" + "hash/crc32" + "testing" +) + +func TestBlockWriterSingleRecord(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + payload := []byte("hello world") + if err := bw.WriteRecord(RecFull, payload, &buf); err != nil { + t.Fatalf("WriteRecord: %v", err) + } + + // Record should still be buffered (block not full). + if buf.Len() != 0 { + t.Fatalf("expected no flush yet, got %d bytes", buf.Len()) + } + + // Flush to get the data. + if err := bw.Flush(&buf); err != nil { + t.Fatalf("Flush: %v", err) + } + + written := buf.Bytes() + + // Verify the record is at the start of a full block. + if len(written) != WalBlockSize { + t.Fatalf("expected full block %d bytes, got %d", WalBlockSize, len(written)) + } + + // Decode and verify the physical record. + rec, consumed, err := DecodePhysicalRecord(written) + if err != nil { + t.Fatalf("DecodePhysicalRecord: %v", err) + } + + if rec.Type != RecFull { + t.Errorf("type = %d, want RecFull(%d)", rec.Type, RecFull) + } + if string(rec.Payload) != "hello world" { + t.Errorf("payload = %q, want %q", rec.Payload, "hello world") + } + + // Remaining bytes after the record should be zero padding. + recEnd := consumed + for i := recEnd; i < WalBlockSize; i++ { + if written[i] != 0 { + t.Errorf("padding byte [%d] = %d, want 0", i, written[i]) + } + } +} + +func TestBlockWriterPadding(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + payloadLen := WalBlockSize - PhysicalRecordHeaderSize + payload := make([]byte, payloadLen) + for i := range payload { + payload[i] = byte(i % 256) + } + + if err := bw.WriteRecord(RecFull, payload, &buf); err != nil { + t.Fatalf("WriteRecord: %v", err) + } + + if buf.Len() != WalBlockSize { + t.Fatalf("expected auto-flush of full block (%d bytes), got %d", WalBlockSize, buf.Len()) + } + + if bw.BlockOffset() != 0 { + t.Errorf("BlockOffset = %d, want 0 after full block write", bw.BlockOffset()) + } + + buf.Reset() + smallPayload := []byte("next") + if err := bw.WriteRecord(RecFull, smallPayload, &buf); err != nil { + t.Fatalf("WriteRecord after full block: %v", err) + } + if buf.Len() != 0 { + t.Fatalf("expected no flush for partial block, got %d bytes", buf.Len()) + } +} + +func TestBlockWriterPaddingNeeded(t *testing.T) { + tests := []struct { + name string + offset uint32 + want int + }{ + {"beginning", 0, 0}, + {"mid_block", 100, 0}, + {"7_remaining", WalBlockSize - 7, 7}, + {"6_remaining", WalBlockSize - 6, 6}, + {"1_remaining", WalBlockSize - 1, 1}, + {"full_block", WalBlockSize, 0}, // would be reset + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + bw := &BlockWriter{offset: tt.offset} + got := bw.paddingNeeded() + if got != tt.want { + t.Errorf("paddingNeeded() = %d, want %d", got, tt.want) + } + }) + } +} + +func TestBlockWriterCrossBlock(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + payloadLen := WalBlockSize - PhysicalRecordHeaderSize - 5 + payload1 := make([]byte, payloadLen) + for i := range payload1 { + payload1[i] = byte('A' + i%26) + } + + if err := bw.WriteRecord(RecFull, payload1, &buf); err != nil { + t.Fatalf("WriteRecord payload1: %v", err) + } + + if buf.Len() != 0 { + t.Fatalf("expected no flush after first record, got %d bytes", buf.Len()) + } + + payload2 := []byte("second") + if err := bw.WriteRecord(RecFull, payload2, &buf); err != nil { + t.Fatalf("WriteRecord payload2: %v", err) + } + + if buf.Len() != WalBlockSize { + t.Fatalf("expected %d bytes flushed, got %d", WalBlockSize, buf.Len()) + } + + firstBlock := buf.Bytes()[:WalBlockSize] + rec, _, err := DecodePhysicalRecord(firstBlock) + if err != nil { + t.Fatalf("DecodePhysicalRecord block 1: %v", err) + } + if rec.Type != RecFull { + t.Errorf("type = %d, want RecFull", rec.Type) + } + if len(rec.Payload) != payloadLen { + t.Errorf("payload len = %d, want %d", len(rec.Payload), payloadLen) + } + for i := WalBlockSize - 5; i < WalBlockSize; i++ { + if firstBlock[i] != 0 { + t.Errorf("padding byte [%d] = %d, want 0", i, firstBlock[i]) + } + } + + if err := bw.Flush(&buf); err != nil { + t.Fatalf("Flush: %v", err) + } + + secondBlock := buf.Bytes()[WalBlockSize:] + if len(secondBlock) != WalBlockSize { + t.Fatalf("second block: expected %d bytes, got %d", WalBlockSize, len(secondBlock)) + } + + rec2, _, err := DecodePhysicalRecord(secondBlock) + if err != nil { + t.Fatalf("DecodePhysicalRecord block 2: %v", err) + } + if string(rec2.Payload) != "second" { + t.Errorf("payload2 = %q, want %q", rec2.Payload, "second") + } +} + +func TestBlockWriterRecordTooLarge(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + // Payload that exceeds block capacity. + payload := make([]byte, WalBlockSize) + err := bw.WriteRecord(RecFull, payload, &buf) + if err == nil { + t.Fatal("expected error for oversized record") + } +} + +func TestBlockWriterMultipleRecords(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + // Write several small records. + records := []struct { + recType uint8 + payload []byte + }{ + {RecFirst, []byte("part1")}, + {RecMiddle, []byte("part2")}, + {RecLast, []byte("part3")}, + } + + for _, r := range records { + if err := bw.WriteRecord(r.recType, r.payload, &buf); err != nil { + t.Fatalf("WriteRecord(%d, %q): %v", r.recType, r.payload, err) + } + } + + if err := bw.Flush(&buf); err != nil { + t.Fatalf("Flush: %v", err) + } + + data := buf.Bytes() + + // Decode all three records from the block. + offset := 0 + for i, expected := range records { + rec, consumed, err := DecodePhysicalRecord(data[offset:]) + if err != nil { + t.Fatalf("record %d: DecodePhysicalRecord at offset %d: %v", i, offset, err) + } + if rec.Type != expected.recType { + t.Errorf("record %d: type = %d, want %d", i, rec.Type, expected.recType) + } + if string(rec.Payload) != string(expected.payload) { + t.Errorf("record %d: payload = %q, want %q", i, rec.Payload, expected.payload) + } + offset += consumed + } +} + +func TestBlockWriterReset(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + if err := bw.WriteRecord(RecFull, []byte("data"), &buf); err != nil { + t.Fatalf("WriteRecord: %v", err) + } + + if bw.BlockOffset() == 0 { + t.Fatal("expected non-zero offset after write") + } + + bw.Reset() + if bw.BlockOffset() != 0 { + t.Errorf("BlockOffset after Reset = %d, want 0", bw.BlockOffset()) + } +} + +func TestBlockWriterFlushEmptyBlock(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + // Flushing an empty block should be a no-op. + if err := bw.Flush(&buf); err != nil { + t.Fatalf("Flush empty: %v", err) + } + if buf.Len() != 0 { + t.Errorf("expected 0 bytes, got %d", buf.Len()) + } +} + +func TestBlockWriterOffsetTracking(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + // Write a small record and verify offset. + payload := []byte("track-me") + if err := bw.WriteRecord(RecFull, payload, &buf); err != nil { + t.Fatalf("WriteRecord: %v", err) + } + + expectedOffset := uint32(PhysicalRecordHeaderSize + len(payload)) + if bw.BlockOffset() != expectedOffset { + t.Errorf("BlockOffset = %d, want %d", bw.BlockOffset(), expectedOffset) + } + + // Flush should write exactly one full block. + if err := bw.Flush(&buf); err != nil { + t.Fatalf("Flush: %v", err) + } + if buf.Len() != WalBlockSize { + t.Errorf("flushed %d bytes, want %d", buf.Len(), WalBlockSize) + } +} + +func TestBlockWriterAutoFlushFullBlock(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + // Fill the block exactly. + payloadLen := WalBlockSize - PhysicalRecordHeaderSize + payload := make([]byte, payloadLen) + for i := range payload { + payload[i] = byte(i) + } + + if err := bw.WriteRecord(RecFull, payload, &buf); err != nil { + t.Fatalf("WriteRecord exact fill: %v", err) + } + + // Block should have been auto-flushed. + if buf.Len() != WalBlockSize { + t.Errorf("expected auto-flush of %d bytes, got %d", WalBlockSize, buf.Len()) + } + if bw.BlockOffset() != 0 { + t.Errorf("BlockOffset after auto-flush = %d, want 0", bw.BlockOffset()) + } + + // Verify CRC is correct by decoding. + data := buf.Bytes() + rec, _, err := DecodePhysicalRecord(data) + if err != nil { + t.Fatalf("DecodePhysicalRecord: %v", err) + } + if len(rec.Payload) != payloadLen { + t.Errorf("payload len = %d, want %d", len(rec.Payload), payloadLen) + } +} + +func TestBlockWriterPhysicalRecordCRC(t *testing.T) { + bw := NewBlockWriter() + var buf bytes.Buffer + + payload := []byte("crc-check") + if err := bw.WriteRecord(RecFull, payload, &buf); err != nil { + t.Fatalf("WriteRecord: %v", err) + } + if err := bw.Flush(&buf); err != nil { + t.Fatalf("Flush: %v", err) + } + + data := buf.Bytes() + + // Manually verify CRC: covers length + type + payload. + storedCRC := binary.LittleEndian.Uint32(data[0:4]) + length := binary.LittleEndian.Uint16(data[4:6]) + recType := data[6] + + if recType != RecFull { + t.Errorf("type = %d, want RecFull", recType) + } + if int(length) != len(payload) { + t.Errorf("length = %d, want %d", length, len(payload)) + } + + // Verify CRC over [length, type, payload]. + crcData := data[4 : 7+length] + computedCRC := crc32.ChecksumIEEE(crcData) + if storedCRC != computedCRC { + t.Errorf("CRC mismatch: stored %d, computed %d", storedCRC, computedCRC) + } +} diff --git a/wal/segment_writer.go b/wal/segment_writer.go new file mode 100644 index 0000000..34287fd --- /dev/null +++ b/wal/segment_writer.go @@ -0,0 +1,170 @@ +package wal + +import ( + "fmt" + "os" + "path/filepath" + + "github.com/dailz/go-kv/config" +) + +// SegmentWriter handles appending WAL batches to a single segment file. +// It manages block-aligned writes via BlockWriter and tracks file offset +// for segment rotation decisions. +type SegmentWriter struct { + fd *os.File + dir string + cfg *config.WalConfig + segmentID uint64 + startSequence uint64 + blockWriter *BlockWriter + currentOffset uint64 // total bytes written (starts at WalFileHeaderSize) + maxPayload uint64 // cfg.MaxSegmentSize - WalFileHeaderSize +} + +// NewSegmentWriter creates a new WAL segment file and writes the file header. +// The segment file is created with a .tmp extension, the header is written and +// synced, then the file is atomically renamed to its final name and synced again. +func NewSegmentWriter( + dir string, + segmentID uint64, + startSequence uint64, + cfg *config.WalConfig, +) (*SegmentWriter, error) { + if err := cfg.Validate(); err != nil { + return nil, fmt.Errorf("wal: invalid config: %w", err) + } + + baseName := fmt.Sprintf("segment-%d.wal", segmentID) + tmpPath := filepath.Join(dir, baseName+".tmp") + finalPath := filepath.Join(dir, baseName) + + // Create the temp file. + fd, err := os.OpenFile(tmpPath, os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0o644) + if err != nil { + return nil, fmt.Errorf("wal: create segment temp file %s: %w", tmpPath, err) + } + + // Build and write the file header. + hdr := &WalFileHeader{ + BlockSize: cfg.BlockSize, + SegmentID: segmentID, + StartSequence: startSequence, + } + encoded := EncodeWalHeader(hdr) + + if _, err := fd.Write(encoded[:]); err != nil { + fd.Close() + os.Remove(tmpPath) + return nil, fmt.Errorf("wal: write segment header: %w", err) + } + + // Sync the header to disk. + if err := fd.Sync(); err != nil { + fd.Close() + os.Remove(tmpPath) + return nil, fmt.Errorf("wal: sync segment header: %w", err) + } + + // Atomically rename temp file to final name. + if err := fd.Close(); err != nil { + os.Remove(tmpPath) + return nil, fmt.Errorf("wal: close temp file: %w", err) + } + + if err := os.Rename(tmpPath, finalPath); err != nil { + os.Remove(tmpPath) + return nil, fmt.Errorf("wal: rename segment file: %w", err) + } + + // Open the final file for appending. + fd, err = os.OpenFile(finalPath, os.O_WRONLY|os.O_APPEND, 0o644) + if err != nil { + return nil, fmt.Errorf("wal: open segment file for append: %w", err) + } + + // Sync directory to make rename durable (best-effort on Linux). + if dirFD, derr := os.Open(dir); derr == nil { + dirFD.Sync() + dirFD.Close() + } + + maxPayload := cfg.MaxSegmentSize - WalFileHeaderSize + + return &SegmentWriter{ + fd: fd, + dir: dir, + cfg: cfg, + segmentID: segmentID, + startSequence: startSequence, + blockWriter: NewBlockWriter(), + currentOffset: WalFileHeaderSize, + maxPayload: maxPayload, + }, nil +} + +// AppendBatch encodes the batch into physical records and appends them to the +// segment file. The encoded batch is split into block-aligned physical records +// using SplitIntoRecords. +func (sw *SegmentWriter) AppendBatch(encodedBatch []byte) error { + records := SplitIntoRecords(encodedBatch) + if len(records) == 0 { + return nil + } + + for _, rec := range records { + if len(rec) < PhysicalRecordHeaderSize { + return fmt.Errorf("wal: corrupted physical record: size %d < header size %d", + len(rec), PhysicalRecordHeaderSize) + } + + recType := rec[6] // type byte is at offset 6 in the encoded record + payload := rec[PhysicalRecordHeaderSize:] + + if err := sw.blockWriter.WriteRecord(recType, payload, sw.fd); err != nil { + return fmt.Errorf("wal: writing physical record: %w", err) + } + + sw.currentOffset += uint64(len(rec)) + } + + return nil +} + +// Sync flushes the segment file to durable storage. +func (sw *SegmentWriter) Sync() error { + return sw.fd.Sync() +} + +// Close flushes any partial block and closes the segment file. +func (sw *SegmentWriter) Close() error { + if err := sw.blockWriter.Flush(sw.fd); err != nil { + return fmt.Errorf("wal: flushing block writer on close: %w", err) + } + return sw.fd.Close() +} + +// RemainingPayload returns the number of bytes that can still be written +// to this segment before it reaches its maximum size. +func (sw *SegmentWriter) RemainingPayload() uint64 { + if sw.currentOffset >= sw.cfg.MaxSegmentSize { + return 0 + } + return sw.cfg.MaxSegmentSize - sw.currentOffset +} + +// CurrentOffset returns the total number of bytes written to the segment file, +// including the file header. +func (sw *SegmentWriter) CurrentOffset() uint64 { + return sw.currentOffset +} + +// SegmentID returns the segment identifier. +func (sw *SegmentWriter) SegmentID() uint64 { + return sw.segmentID +} + +// SegmentPath returns the full filesystem path to the segment file. +func (sw *SegmentWriter) SegmentPath() string { + return filepath.Join(sw.dir, fmt.Sprintf("segment-%d.wal", sw.segmentID)) +} diff --git a/wal/segment_writer_test.go b/wal/segment_writer_test.go new file mode 100644 index 0000000..2423b80 --- /dev/null +++ b/wal/segment_writer_test.go @@ -0,0 +1,339 @@ +package wal + +import ( + "os" + "path/filepath" + "testing" + + "github.com/dailz/go-kv/config" +) + +func testWalConfig() *config.WalConfig { + cfg := config.Defaults() + return &cfg +} + +func TestSegmentWriterCreation(t *testing.T) { + dir := t.TempDir() + cfg := testWalConfig() + + sw, err := NewSegmentWriter(dir, 1, 100, cfg) + if err != nil { + t.Fatalf("NewSegmentWriter: %v", err) + } + + expectedPath := filepath.Join(dir, "segment-1.wal") + if sw.SegmentPath() != expectedPath { + t.Errorf("SegmentPath = %q, want %q", sw.SegmentPath(), expectedPath) + } + if sw.SegmentID() != 1 { + t.Errorf("SegmentID = %d, want 1", sw.SegmentID()) + } + if sw.CurrentOffset() != WalFileHeaderSize { + t.Errorf("CurrentOffset = %d, want %d", sw.CurrentOffset(), WalFileHeaderSize) + } + + // Verify the file exists and has correct header. + data, err := os.ReadFile(expectedPath) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + if len(data) != WalFileHeaderSize { + t.Errorf("file size = %d, want %d (header only)", len(data), WalFileHeaderSize) + } + + hdr, err := DecodeWalHeader(data) + if err != nil { + t.Fatalf("DecodeWalHeader: %v", err) + } + if hdr.SegmentID != 1 { + t.Errorf("header SegmentID = %d, want 1", hdr.SegmentID) + } + if hdr.StartSequence != 100 { + t.Errorf("header StartSequence = %d, want 100", hdr.StartSequence) + } + if hdr.BlockSize != cfg.BlockSize { + t.Errorf("header BlockSize = %d, want %d", hdr.BlockSize, cfg.BlockSize) + } + + // No .tmp file should remain. + tmpPath := filepath.Join(dir, "segment-1.wal.tmp") + if _, err := os.Stat(tmpPath); !os.IsNotExist(err) { + t.Errorf("temp file %q should not exist", tmpPath) + } + + if err := sw.Close(); err != nil { + t.Fatalf("Close: %v", err) + } +} + +func TestSegmentWriterAppendBatch(t *testing.T) { + dir := t.TempDir() + cfg := testWalConfig() + + sw, err := NewSegmentWriter(dir, 42, 0, cfg) + if err != nil { + t.Fatalf("NewSegmentWriter: %v", err) + } + t.Cleanup(func() { sw.Close() }) + + // Encode a small batch. + entries := []*WalEntry{ + {OpType: OpPut, ValueKind: VKInline, Key: []byte("key1"), Value: []byte("val1")}, + {OpType: OpPut, ValueKind: VKInline, Key: []byte("key2"), Value: []byte("val2")}, + } + encoded, err := EncodeWalBatch(0, entries) + if err != nil { + t.Fatalf("EncodeWalBatch: %v", err) + } + + if err := sw.AppendBatch(encoded); err != nil { + t.Fatalf("AppendBatch: %v", err) + } + + if err := sw.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + // Read back the file and verify records. + data, err := os.ReadFile(sw.SegmentPath()) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + + // Skip file header. + body := data[WalFileHeaderSize:] + + // Use FragmentCollector to reassemble. + fc := NewFragmentCollector() + offset := 0 + for offset < len(body) { + // Check for trailing zeros (block padding). + if body[offset] == 0 { + break + } + + rec, consumed, err := DecodePhysicalRecord(body[offset:]) + if err != nil { + t.Fatalf("DecodePhysicalRecord at offset %d: %v", offset, err) + } + + if err := fc.Append(rec.Type, rec.Payload); err != nil { + t.Fatalf("FragmentCollector.Append: %v", err) + } + offset += consumed + } + + if !fc.IsComplete() { + t.Fatal("fragment collector should be complete") + } + + decoded, err := DecodeWalBatch(fc.BatchData()) + if err != nil { + t.Fatalf("DecodeWalBatch: %v", err) + } + + if decoded.EntryCount != 2 { + t.Errorf("EntryCount = %d, want 2", decoded.EntryCount) + } + if decoded.BaseSequence != 0 { + t.Errorf("BaseSequence = %d, want 0", decoded.BaseSequence) + } +} + +func TestSegmentWriterMultipleBatches(t *testing.T) { + dir := t.TempDir() + cfg := testWalConfig() + + sw, err := NewSegmentWriter(dir, 1, 0, cfg) + if err != nil { + t.Fatalf("NewSegmentWriter: %v", err) + } + + for i := 0; i < 5; i++ { + entries := []*WalEntry{ + { + OpType: OpPut, + ValueKind: VKInline, + Key: []byte("key"), + Value: []byte("val"), + }, + } + encoded, err := EncodeWalBatch(uint64(i), entries) + if err != nil { + t.Fatalf("EncodeWalBatch %d: %v", i, err) + } + if err := sw.AppendBatch(encoded); err != nil { + t.Fatalf("AppendBatch %d: %v", i, err) + } + } + + if err := sw.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + data, err := os.ReadFile(sw.SegmentPath()) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + + body := data[WalFileHeaderSize:] + fc := NewFragmentCollector() + batchCount := 0 + offset := 0 + + for offset < len(body) { + if body[offset] == 0 { + break + } + + rec, consumed, err := DecodePhysicalRecord(body[offset:]) + if err != nil { + t.Fatalf("DecodePhysicalRecord at offset %d: %v", offset, err) + } + offset += consumed + + if err := fc.Append(rec.Type, rec.Payload); err != nil { + t.Fatalf("FragmentCollector.Append: %v", err) + } + + if fc.IsComplete() { + batchCount++ + decoded, err := DecodeWalBatch(fc.BatchData()) + if err != nil { + t.Fatalf("DecodeWalBatch %d: %v", batchCount, err) + } + if decoded.BaseSequence != uint64(batchCount-1) { + t.Errorf("batch %d BaseSequence = %d, want %d", + batchCount, decoded.BaseSequence, batchCount-1) + } + fc.Reset() + } + } + + if batchCount != 5 { + t.Errorf("decoded %d batches, want 5", batchCount) + } +} + +func TestSegmentWriterRemainingPayload(t *testing.T) { + dir := t.TempDir() + cfg := testWalConfig() + + sw, err := NewSegmentWriter(dir, 1, 0, cfg) + if err != nil { + t.Fatalf("NewSegmentWriter: %v", err) + } + t.Cleanup(func() { sw.Close() }) + + initialRemaining := sw.RemainingPayload() + + entries := []*WalEntry{ + {OpType: OpPut, ValueKind: VKInline, Key: []byte("k"), Value: []byte("v")}, + } + encoded, err := EncodeWalBatch(0, entries) + if err != nil { + t.Fatalf("EncodeWalBatch: %v", err) + } + if err := sw.AppendBatch(encoded); err != nil { + t.Fatalf("AppendBatch: %v", err) + } + + if sw.RemainingPayload() >= initialRemaining { + t.Errorf("RemainingPayload should decrease after write, got %d >= %d", + sw.RemainingPayload(), initialRemaining) + } +} + +func TestSegmentWriterLargeBatch(t *testing.T) { + dir := t.TempDir() + cfg := testWalConfig() + + sw, err := NewSegmentWriter(dir, 1, 0, cfg) + if err != nil { + t.Fatalf("NewSegmentWriter: %v", err) + } + + // Create a batch larger than one block payload (~32KB - 7 bytes). + largeValue := make([]byte, 40*1024) + for i := range largeValue { + largeValue[i] = byte(i % 256) + } + + entries := []*WalEntry{ + {OpType: OpPut, ValueKind: VKValueLogPointer, Key: []byte("large-key"), Value: largeValue}, + } + encoded, err := EncodeWalBatch(0, entries) + if err != nil { + t.Fatalf("EncodeWalBatch: %v", err) + } + + if err := sw.AppendBatch(encoded); err != nil { + t.Fatalf("AppendBatch: %v", err) + } + + if err := sw.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + + // Read back and verify. + data, err := os.ReadFile(sw.SegmentPath()) + if err != nil { + t.Fatalf("ReadFile: %v", err) + } + + body := data[WalFileHeaderSize:] + fc := NewFragmentCollector() + offset := 0 + fragCount := 0 + + for offset < len(body) { + if body[offset] == 0 { + break + } + + rec, consumed, err := DecodePhysicalRecord(body[offset:]) + if err != nil { + t.Fatalf("DecodePhysicalRecord at offset %d: %v", offset, err) + } + offset += consumed + fragCount++ + + if err := fc.Append(rec.Type, rec.Payload); err != nil { + t.Fatalf("FragmentCollector.Append type=%d: %v", rec.Type, err) + } + } + + if !fc.IsComplete() { + t.Fatal("fragment collector should be complete after large batch") + } + if fragCount < 2 { + t.Errorf("expected multiple fragments for large batch, got %d", fragCount) + } + + decoded, err := DecodeWalBatch(fc.BatchData()) + if err != nil { + t.Fatalf("DecodeWalBatch: %v", err) + } + if decoded.EntryCount != 1 { + t.Errorf("EntryCount = %d, want 1", decoded.EntryCount) + } +} + +func TestSegmentWriterSync(t *testing.T) { + dir := t.TempDir() + cfg := testWalConfig() + + sw, err := NewSegmentWriter(dir, 1, 0, cfg) + if err != nil { + t.Fatalf("NewSegmentWriter: %v", err) + } + + if err := sw.Sync(); err != nil { + t.Fatalf("Sync: %v", err) + } + + if err := sw.Close(); err != nil { + t.Fatalf("Close: %v", err) + } +}