diff --git a/.omo/boulder.json b/.omo/boulder.json index 2fdcb88..117db6a 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:49:50.501Z", + "updated_at": "2026-06-12T05:57:16.412Z", "session_ids": [ "opencode:ses_145c3bae9ffeTB2zbsTym0Cev8" ], @@ -94,6 +94,19 @@ "status": "completed", "ended_at": "2026-06-12T05:49:50.501Z", "elapsed_ms": 106010 + }, + "todo:13": { + "task_key": "todo:13", + "task_label": "13", + "task_title": "Commit queue + WAL writer main loop", + "session_id": "opencode:ses_1459dbbd5ffewqrKSfu5f6t38q", + "agent": "Sisyphus-Junior", + "category": "deep", + "updated_at": "2026-06-12T05:57:16.412Z", + "started_at": "2026-06-12T05:56:53.263Z", + "status": "completed", + "ended_at": "2026-06-12T05:57:16.412Z", + "elapsed_ms": 23149 } } } @@ -101,7 +114,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:49:50.501Z", + "updated_at": "2026-06-12T05:57:16.412Z", "session_ids": [ "opencode:ses_145c3bae9ffeTB2zbsTym0Cev8" ], @@ -187,6 +200,19 @@ "status": "completed", "ended_at": "2026-06-12T05:49:50.501Z", "elapsed_ms": 106010 + }, + "todo:13": { + "task_key": "todo:13", + "task_label": "13", + "task_title": "Commit queue + WAL writer main loop", + "session_id": "opencode:ses_1459dbbd5ffewqrKSfu5f6t38q", + "agent": "Sisyphus-Junior", + "category": "deep", + "updated_at": "2026-06-12T05:57:16.412Z", + "started_at": "2026-06-12T05:56:53.263Z", + "status": "completed", + "ended_at": "2026-06-12T05:57:16.412Z", + "elapsed_ms": 23149 } }, "agent": "atlas" diff --git a/.omo/plans/phase1-wal.md b/.omo/plans/phase1-wal.md index 32d76cd..cc3e8cd 100644 --- a/.omo/plans/phase1-wal.md +++ b/.omo/plans/phase1-wal.md @@ -1009,7 +1009,7 @@ Max Concurrent: 6 (Wave 1b) **Commit**: YES (with Task 13) -- [ ] 13. Commit queue + WAL writer main loop +- [x] 13. Commit queue + WAL writer main loop **What to do**: - 创建 `wal/commit_queue.go`:`CommitQueue` — 写请求队列 @@ -1233,7 +1233,7 @@ Max Concurrent: 6 (Wave 1b) **Commit**: YES (with Tasks 16, 17) -- [ ] 16. WAL recovery: fragment collector + batch replay +- [x] 16. WAL recovery: fragment collector + batch replay **What to do**: - 创建 `wal/recovery.go`:恢复核心逻辑 diff --git a/wal/commit_queue.go b/wal/commit_queue.go new file mode 100644 index 0000000..bcae13f --- /dev/null +++ b/wal/commit_queue.go @@ -0,0 +1,58 @@ +package wal + +// WriteResult is sent to the caller when the write completes or fails. +type WriteResult struct { + Sequence uint64 + Err error +} + +// CommitRequest represents a single write request submitted by a caller. +type CommitRequest struct { + Entries []*WalEntry + Result chan WriteResult +} + +// CommitQueue is a bounded buffered channel for write requests. +type CommitQueue struct { + ch chan *CommitRequest +} + +// NewCommitQueue creates a bounded queue for commit requests. +func NewCommitQueue(capacity int) *CommitQueue { + if capacity < 1 { + capacity = 1 + } + return &CommitQueue{ch: make(chan *CommitRequest, capacity)} +} + +// Submit creates a commit request, submits it to the queue, and returns it. +// The send blocks when the queue is full, preserving backpressure. +func (cq *CommitQueue) Submit(entries []*WalEntry) *CommitRequest { + req := &CommitRequest{ + Entries: entries, + Result: make(chan WriteResult, 1), + } + cq.ch <- req + return req +} + +// Collect drains all currently pending requests from the queue. +func (cq *CommitQueue) Collect() []*CommitRequest { + var requests []*CommitRequest + for { + select { + case req, ok := <-cq.ch: + if !ok { + return requests + } + requests = append(requests, req) + default: + return requests + } + } +} + +// Close closes the queue. Callers must stop submitting before Close. +func (cq *CommitQueue) Close() { + close(cq.ch) +} diff --git a/wal/recovery.go b/wal/recovery.go new file mode 100644 index 0000000..324d2ab --- /dev/null +++ b/wal/recovery.go @@ -0,0 +1,169 @@ +package wal + +import ( + "errors" + "fmt" + "math" +) + +// BatchReplayer is the interface for replaying recovered WAL entries. +// The recovery process calls the appropriate method for each entry. +type BatchReplayer interface { + ReplayPut(key, value []byte, sequence uint64) + ReplayDelete(key []byte, sequence uint64) +} + +type replayAction struct { + opType uint8 + key []byte + value []byte + sequence uint64 +} + +// ReplayBatch validates and replays a decoded WAL batch, returning the next +// expected sequence after the batch. +func ReplayBatch(batch *WalBatch, expectedSequence uint64, replayer BatchReplayer) (nextSequence uint64, err error) { + if batch == nil { + return expectedSequence, errors.New("wal: batch is nil") + } + if replayer == nil { + return expectedSequence, errors.New("wal: batch replayer is nil") + } + + if batch.Flags != 0 { + return expectedSequence, fmt.Errorf("wal: invalid batch flags %d", batch.Flags) + } + if batch.EntryCount == 0 { + return expectedSequence, errors.New("wal: batch entry count is zero") + } + if batch.EntryCount > MaxWalBatchEntryCount { + return expectedSequence, fmt.Errorf("wal: batch entry count %d exceeds maximum %d", batch.EntryCount, MaxWalBatchEntryCount) + } + if batch.EntriesSize != uint32(len(batch.Entries)) { + return expectedSequence, fmt.Errorf("wal: batch entries size mismatch: header says %d bytes, have %d bytes", batch.EntriesSize, len(batch.Entries)) + } + if batch.EntriesSize == 0 { + return expectedSequence, errors.New("wal: batch entries size is zero") + } + if batch.EntriesSize > MaxWalBatchEntriesSize { + return expectedSequence, fmt.Errorf("wal: batch entries size %d exceeds maximum %d", batch.EntriesSize, MaxWalBatchEntriesSize) + } + if batch.BaseSequence != expectedSequence { + return expectedSequence, fmt.Errorf("wal: batch base sequence %d does not match expected sequence %d", batch.BaseSequence, expectedSequence) + } + + entryCount := uint64(batch.EntryCount) + if batch.BaseSequence > math.MaxUint64-(entryCount-1) { + return expectedSequence, fmt.Errorf("wal: batch sequence range overflows uint64: base sequence %d entry count %d", batch.BaseSequence, batch.EntryCount) + } + if expectedSequence > math.MaxUint64-entryCount { + return expectedSequence, fmt.Errorf("wal: next sequence overflows uint64: expected sequence %d entry count %d", expectedSequence, batch.EntryCount) + } + + actions := make([]replayAction, 0, batch.EntryCount) + offset := 0 + for i := range entryCount { + entry, consumed, err := DecodeEntry(batch.Entries[offset:]) + if err != nil { + return expectedSequence, fmt.Errorf("wal: decode batch entry: %w", err) + } + if consumed <= 0 { + return expectedSequence, errors.New("wal: decoded batch entry consumed no bytes") + } + + sequence := batch.BaseSequence + i + actions = append(actions, replayAction{ + opType: entry.OpType, + key: entry.Key, + value: entry.Value, + sequence: sequence, + }) + offset += consumed + } + + if offset != len(batch.Entries) { + return expectedSequence, fmt.Errorf("wal: batch entries contain trailing bytes: parsed %d of %d", offset, len(batch.Entries)) + } + + for _, action := range actions { + switch action.opType { + case OpPut: + replayer.ReplayPut(action.key, action.value, action.sequence) + case OpDelete: + replayer.ReplayDelete(action.key, action.sequence) + default: + return expectedSequence, fmt.Errorf("wal: invalid op type %d", action.opType) + } + } + + return expectedSequence + entryCount, nil +} + +// ReplaySegmentFile replays all complete WAL batches from one segment file. +func ReplaySegmentFile(filePath string, startSequence uint64, replayer BatchReplayer) (nextSequence uint64, err error) { + nextSequence = startSequence + records, parseErr := ParseRecordsFromFile(filePath) + if parseErr != nil && !IsTailCorruption(parseErr) { + return nextSequence, fmt.Errorf("wal: parse segment records: %w", parseErr) + } + + collector := NewFragmentCollector() + for _, record := range records { + if err := collector.Append(record.Type, record.Payload); err != nil { + return nextSequence, fmt.Errorf("wal: collect segment fragments: %w", err) + } + if !collector.IsComplete() { + continue + } + + batch, err := DecodeWalBatch(collector.BatchData()) + if err != nil { + return nextSequence, fmt.Errorf("wal: decode recovered batch: %w", err) + } + nextSequence, err = ReplayBatch(batch, nextSequence, replayer) + if err != nil { + return nextSequence, fmt.Errorf("wal: replay recovered batch: %w", err) + } + collector.Reset() + } + + if parseErr != nil { + return nextSequence, parseErr + } + if collector.State() == FragmentCollecting { + return nextSequence, &TailCorruptionError{ + Offset: 0, + Err: errors.New("incomplete fragmented batch at segment tail"), + } + } + + return nextSequence, nil +} + +// RecoverFromSegments scans and replays WAL segments from recoverySegmentID. +func RecoverFromSegments(dir string, recoverySegmentID uint64, replayer BatchReplayer) (nextSequence uint64, err error) { + segments, err := ScanSegments(dir, recoverySegmentID) + if err != nil { + return 0, fmt.Errorf("wal: scan recovery segments: %w", err) + } + if len(segments) == 0 { + return 0, nil + } + + nextSequence = segments[0].StartSequence + for _, segment := range segments { + if segment.StartSequence != nextSequence { + return nextSequence, fmt.Errorf("wal: segment start sequence %d does not match expected sequence %d", segment.StartSequence, nextSequence) + } + + nextSequence, err = ReplaySegmentFile(segment.FilePath, nextSequence, replayer) + if err != nil { + if IsTailCorruption(err) { + return nextSequence, err + } + return nextSequence, fmt.Errorf("wal: replay segment: %w", err) + } + } + + return nextSequence, nil +} diff --git a/wal/recovery_test.go b/wal/recovery_test.go new file mode 100644 index 0000000..134532a --- /dev/null +++ b/wal/recovery_test.go @@ -0,0 +1,252 @@ +package wal + +import ( + "errors" + "math" + "os" + "reflect" + "strings" + "testing" + + "github.com/dailz/go-kv/config" +) + +type mockReplayer struct { + puts []replayPut + deletes []replayDelete +} + +type replayPut struct { + key string + value string + seq uint64 +} + +type replayDelete struct { + key string + seq uint64 +} + +func (m *mockReplayer) ReplayPut(key, value []byte, sequence uint64) { + m.puts = append(m.puts, replayPut{key: string(key), value: string(value), seq: sequence}) +} + +func (m *mockReplayer) ReplayDelete(key []byte, sequence uint64) { + m.deletes = append(m.deletes, replayDelete{key: string(key), seq: sequence}) +} + +func TestReplayBatchValid(t *testing.T) { + encoded, err := EncodeWalBatch(100, []*WalEntry{ + makePutEntry("alpha", "one"), + makeDeleteEntry("beta"), + }) + if err != nil { + t.Fatalf("EncodeWalBatch: %v", err) + } + batch, err := DecodeWalBatch(encoded) + if err != nil { + t.Fatalf("DecodeWalBatch: %v", err) + } + + replayer := &mockReplayer{} + next, err := ReplayBatch(batch, 100, replayer) + if err != nil { + t.Fatalf("ReplayBatch: %v", err) + } + if next != 102 { + t.Fatalf("nextSequence = %d, want 102", next) + } + + wantPuts := []replayPut{{key: "alpha", value: "one", seq: 100}} + if !reflect.DeepEqual(replayer.puts, wantPuts) { + t.Fatalf("puts = %#v, want %#v", replayer.puts, wantPuts) + } + wantDeletes := []replayDelete{{key: "beta", seq: 101}} + if !reflect.DeepEqual(replayer.deletes, wantDeletes) { + t.Fatalf("deletes = %#v, want %#v", replayer.deletes, wantDeletes) + } +} + +func TestReplayBatchSequenceMismatch(t *testing.T) { + batch := mustDecodeTestBatch(t, 10, []*WalEntry{makePutEntry("k", "v")}) + + _, err := ReplayBatch(batch, 11, &mockReplayer{}) + if err == nil { + t.Fatal("ReplayBatch succeeded, want sequence mismatch error") + } + if !strings.Contains(err.Error(), "base sequence") { + t.Fatalf("error = %v, want base sequence context", err) + } +} + +func TestReplayBatchZeroEntries(t *testing.T) { + batch := &WalBatch{ + Flags: 0, + BaseSequence: 10, + EntryCount: 0, + EntriesSize: 0, + Entries: nil, + } + + _, err := ReplayBatch(batch, 10, &mockReplayer{}) + if err == nil { + t.Fatal("ReplayBatch succeeded, want zero entries error") + } +} + +func TestReplayBatchOverflowCheck(t *testing.T) { + entryBytes, err := EncodeEntry(makePutEntry("k", "v")) + if err != nil { + t.Fatalf("EncodeEntry: %v", err) + } + batch := &WalBatch{ + Flags: 0, + BaseSequence: math.MaxUint64 - 1, + EntryCount: 3, + EntriesSize: uint32(len(entryBytes) * 3), + Entries: append(append(append([]byte{}, entryBytes...), entryBytes...), entryBytes...), + } + + _, err = ReplayBatch(batch, math.MaxUint64-1, &mockReplayer{}) + if err == nil { + t.Fatal("ReplayBatch succeeded, want overflow error") + } +} + +func TestReplaySegmentFileFull(t *testing.T) { + dir := t.TempDir() + filePath := writeTestSegment(t, dir, 0, 50, [][]*WalEntry{ + {makePutEntry("a", "1"), makeDeleteEntry("b")}, + {makePutEntry("c", "3")}, + }) + + replayer := &mockReplayer{} + next, err := ReplaySegmentFile(filePath, 50, replayer) + if err != nil { + t.Fatalf("ReplaySegmentFile: %v", err) + } + if next != 53 { + t.Fatalf("nextSequence = %d, want 53", next) + } + + wantPuts := []replayPut{{key: "a", value: "1", seq: 50}, {key: "c", value: "3", seq: 52}} + if !reflect.DeepEqual(replayer.puts, wantPuts) { + t.Fatalf("puts = %#v, want %#v", replayer.puts, wantPuts) + } + wantDeletes := []replayDelete{{key: "b", seq: 51}} + if !reflect.DeepEqual(replayer.deletes, wantDeletes) { + t.Fatalf("deletes = %#v, want %#v", replayer.deletes, wantDeletes) + } +} + +func TestReplaySegmentFileTailCorruption(t *testing.T) { + dir := t.TempDir() + filePath := writeTestSegment(t, dir, 0, 70, [][]*WalEntry{ + {makePutEntry("ok", "before-corruption")}, + }) + appendFileBytes(t, filePath, []byte{0x01, 0x02, 0x03}) + + replayer := &mockReplayer{} + next, err := ReplaySegmentFile(filePath, 70, replayer) + if err == nil { + t.Fatal("ReplaySegmentFile succeeded, want tail corruption error") + } + if !IsTailCorruption(err) { + t.Fatalf("error = %v, want tail corruption", err) + } + if next != 71 { + t.Fatalf("nextSequence = %d, want 71", next) + } + wantPuts := []replayPut{{key: "ok", value: "before-corruption", seq: 70}} + if !reflect.DeepEqual(replayer.puts, wantPuts) { + t.Fatalf("puts = %#v, want %#v", replayer.puts, wantPuts) + } +} + +func TestRecoverFromSegmentsMultiple(t *testing.T) { + dir := t.TempDir() + writeTestSegment(t, dir, 0, 10, [][]*WalEntry{ + {makePutEntry("s0-a", "a"), makeDeleteEntry("s0-b")}, + }) + writeTestSegment(t, dir, 1, 12, [][]*WalEntry{ + {makePutEntry("s1-c", "c")}, + {makeDeleteEntry("s1-d")}, + }) + + replayer := &mockReplayer{} + next, err := RecoverFromSegments(dir, 0, replayer) + if err != nil { + t.Fatalf("RecoverFromSegments: %v", err) + } + if next != 14 { + t.Fatalf("nextSequence = %d, want 14", next) + } + + wantPuts := []replayPut{{key: "s0-a", value: "a", seq: 10}, {key: "s1-c", value: "c", seq: 12}} + if !reflect.DeepEqual(replayer.puts, wantPuts) { + t.Fatalf("puts = %#v, want %#v", replayer.puts, wantPuts) + } + wantDeletes := []replayDelete{{key: "s0-b", seq: 11}, {key: "s1-d", seq: 13}} + if !reflect.DeepEqual(replayer.deletes, wantDeletes) { + t.Fatalf("deletes = %#v, want %#v", replayer.deletes, wantDeletes) + } +} + +func mustDecodeTestBatch(t *testing.T, baseSequence uint64, entries []*WalEntry) *WalBatch { + t.Helper() + encoded, err := EncodeWalBatch(baseSequence, entries) + if err != nil { + t.Fatalf("EncodeWalBatch: %v", err) + } + batch, err := DecodeWalBatch(encoded) + if err != nil { + t.Fatalf("DecodeWalBatch: %v", err) + } + return batch +} + +func writeTestSegment(t *testing.T, dir string, segmentID, startSequence uint64, batches [][]*WalEntry) string { + t.Helper() + cfg := config.Defaults() + sw, err := NewSegmentWriter(dir, segmentID, startSequence, &cfg) + if err != nil { + t.Fatalf("NewSegmentWriter: %v", err) + } + closed := false + defer func() { + if !closed { + if closeErr := sw.Close(); closeErr != nil && !errors.Is(closeErr, os.ErrClosed) { + t.Fatalf("SegmentWriter.Close cleanup: %v", closeErr) + } + } + }() + + nextSequence := startSequence + for i, entries := range batches { + encoded, err := EncodeWalBatch(nextSequence, 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) + } + nextSequence += uint64(len(entries)) + } + if err := sw.Close(); err != nil { + t.Fatalf("SegmentWriter.Close: %v", err) + } + closed = true + return sw.SegmentPath() +} + +func appendFileBytes(t *testing.T, filePath string, data []byte) { + t.Helper() + f, err := os.OpenFile(filePath, os.O_WRONLY|os.O_APPEND, 0) + if err != nil { + t.Fatalf("OpenFile append: %v", err) + } + defer f.Close() + if _, err := f.Write(data); err != nil { + t.Fatalf("Write corruption bytes: %v", err) + } +} diff --git a/wal/writer.go b/wal/writer.go new file mode 100644 index 0000000..7de8261 --- /dev/null +++ b/wal/writer.go @@ -0,0 +1,342 @@ +package wal + +import ( + "errors" + "sync" + "sync/atomic" + "time" + + go_kv "github.com/dailz/go-kv" + "github.com/dailz/go-kv/config" + "github.com/dailz/go-kv/memtable" +) + +// MemTableList tracks the active memtable plus frozen memtables retained for reads. +type MemTableList struct { + active *memtable.MemTable + immutable []*memtable.MemTable + mu sync.RWMutex +} + +// GetActive returns the current writable memtable. +func (mtl *MemTableList) GetActive() *memtable.MemTable { + mtl.mu.RLock() + defer mtl.mu.RUnlock() + return mtl.active +} + +// GetAll returns memtables from newest to oldest for read lookup. +func (mtl *MemTableList) GetAll() []*memtable.MemTable { + mtl.mu.RLock() + defer mtl.mu.RUnlock() + + out := make([]*memtable.MemTable, 0, 1+len(mtl.immutable)) + if mtl.active != nil { + out = append(out, mtl.active) + } + for i := len(mtl.immutable) - 1; i >= 0; i-- { + out = append(out, mtl.immutable[i]) + } + return out +} + +func (mtl *MemTableList) rotateActive(newActive *memtable.MemTable) { + mtl.mu.Lock() + defer mtl.mu.Unlock() + if mtl.active != nil { + mtl.immutable = append(mtl.immutable, mtl.active) + } + mtl.active = newActive +} + +// WalWriter owns the single WAL append goroutine and group-commit pipeline. +type WalWriter struct { + cfg *config.WalConfig + dir string + queue *CommitQueue + segManager *SegmentManager + seqManager *SequenceManager + memTables *MemTableList + writeStopped atomic.Bool + done chan struct{} + wg sync.WaitGroup + closeOnce sync.Once + closeErr error +} + +// NewWalWriter creates a WAL writer and starts its single background loop. +func NewWalWriter(cfg *config.WalConfig, dir string, startSegmentID uint64, startSequence uint64) (*WalWriter, error) { + if cfg == nil { + defaults := config.Defaults() + cfg = &defaults + } + if err := cfg.Validate(); err != nil { + return nil, err + } + + segManager, err := NewSegmentManager(dir, startSegmentID, startSequence, cfg) + if err != nil { + return nil, err + } + + queueCapacity := max(1, int(cfg.MaxBatchEntries)) + + ww := &WalWriter{ + cfg: cfg, + dir: dir, + queue: NewCommitQueue(queueCapacity), + segManager: segManager, + seqManager: NewSequenceManager(startSequence), + memTables: &MemTableList{ + active: memtable.NewMemTable(cfg.MemTableSize), + }, + done: make(chan struct{}), + } + + ww.wg.Add(1) + go ww.runLoop() + + return ww, nil +} + +func (ww *WalWriter) runLoop() { + defer ww.wg.Done() + for { + select { + case <-ww.done: + ww.processRemaining() + return + default: + } + + requests := ww.queue.Collect() + if len(requests) == 0 { + time.Sleep(100 * time.Microsecond) + continue + } + + ww.processBatch(requests) + } +} + +func (ww *WalWriter) processRemaining() { + for { + requests := ww.queue.Collect() + if len(requests) == 0 { + return + } + ww.processBatch(requests) + } +} + +func (ww *WalWriter) processBatch(requests []*CommitRequest) { + if ww.writeStopped.Load() { + ww.sendError(requests, go_kv.ErrWriteStopped) + return + } + + entries, requestOffsets := flattenRequests(requests) + if err := ValidateBatchLimits(entries, ww.cfg); err != nil { + ww.sendError(requests, err) + return + } + if err := ww.reserveMemTable(entries); err != nil { + ww.sendError(requests, err) + return + } + + baseSequence, err := ww.seqManager.AllocateBatch(uint32(len(entries))) + if err != nil { + ww.sendError(requests, err) + return + } + lastSequence := baseSequence + uint64(len(entries)) - 1 + + encoded, err := EncodeWalBatch(baseSequence, entries) + if err != nil { + ww.stopWithError(requests, err) + return + } + + if err := ww.segManager.AppendBatch(encoded); err != nil { + ww.stopWithError(requests, go_kv.ErrCommitUnknown) + return + } + + active := ww.memTables.GetActive() + if err := writePending(active, entries, baseSequence); err != nil { + abortRange(active, baseSequence, lastSequence) + ww.stopWithError(requests, go_kv.ErrCommitUnknown) + return + } + + if err := ww.segManager.Sync(); err != nil { + abortRange(active, baseSequence, lastSequence) + ww.stopWithError(requests, go_kv.ErrCommitUnknown) + return + } + + ww.seqManager.MarkDurable(lastSequence) + active.Publish(lastSequence) + ww.seqManager.Publish(lastSequence) + ww.sendSuccess(requests, baseSequence, requestOffsets) +} + +func flattenRequests(requests []*CommitRequest) ([]*WalEntry, []uint64) { + var entries []*WalEntry + requestOffsets := make([]uint64, len(requests)) + for i, req := range requests { + requestOffsets[i] = uint64(len(entries)) + entries = append(entries, req.Entries...) + } + return entries, requestOffsets +} + +func (ww *WalWriter) reserveMemTable(entries []*WalEntry) error { + reserveEntries := make([]memtable.ReserveEntry, 0, len(entries)) + for _, entry := range entries { + reserveEntries = append(reserveEntries, memtable.ReserveEntry{ + Key: entry.Key, + Value: entry.Value, + IsDelete: entry.OpType == OpDelete, + }) + } + + active := ww.memTables.GetActive() + if _, err := active.Reserve(reserveEntries); err == nil { + return nil + } else if !errors.Is(err, memtable.ErrMemTableFull) { + return err + } + + ww.memTables.rotateActive(memtable.NewMemTable(ww.cfg.MemTableSize)) + active = ww.memTables.GetActive() + _, err := active.Reserve(reserveEntries) + return err +} + +func writePending(mt *memtable.MemTable, entries []*WalEntry, baseSequence uint64) error { + for i, entry := range entries { + sequence := baseSequence + uint64(i) + switch entry.OpType { + case OpPut: + if err := mt.PutPending(entry.Key, entry.Value, sequence); err != nil { + return err + } + case OpDelete: + if err := mt.DeletePending(entry.Key, sequence); err != nil { + return err + } + default: + return errors.New("wal: unsupported operation in committed batch") + } + } + return nil +} + +func abortRange(mt *memtable.MemTable, firstSequence uint64, lastSequence uint64) { + for seq := firstSequence; seq <= lastSequence; seq++ { + mt.Abort(seq) + if seq == ^uint64(0) { + return + } + } +} + +func (ww *WalWriter) stopWithError(requests []*CommitRequest, err error) { + ww.writeStopped.Store(true) + ww.sendError(requests, err) +} + +func (ww *WalWriter) sendError(requests []*CommitRequest, err error) { + for _, req := range requests { + req.Result <- WriteResult{Err: err} + } +} + +func (ww *WalWriter) sendSuccess(requests []*CommitRequest, baseSequence uint64, requestOffsets []uint64) { + for i, req := range requests { + req.Result <- WriteResult{Sequence: baseSequence + requestOffsets[i]} + } +} + +// Put stores key with value. +func (ww *WalWriter) Put(key, value []byte) error { + if ww.writeStopped.Load() { + return go_kv.ErrWriteStopped + } + + req := ww.queue.Submit([]*WalEntry{{ + OpType: OpPut, + ValueKind: VKInline, + Key: cloneBytes(key), + Value: cloneBytes(value), + }}) + result := <-req.Result + return result.Err +} + +// Delete removes key. +func (ww *WalWriter) Delete(key []byte) error { + if ww.writeStopped.Load() { + return go_kv.ErrWriteStopped + } + + req := ww.queue.Submit([]*WalEntry{{ + OpType: OpDelete, + ValueKind: VKNone, + Key: cloneBytes(key), + }}) + result := <-req.Result + return result.Err +} + +// Get returns the newest visible value across active and immutable memtables. +func (ww *WalWriter) Get(key []byte) *memtable.GetResult { + for _, mt := range ww.memTables.GetAll() { + result := mt.Get(key) + if !result.Found { + continue + } + if result.Value == nil { + return &memtable.GetResult{Found: false, Sequence: result.Sequence} + } + return result + } + return &memtable.GetResult{Found: false} +} + +// GetDurableSequence returns the highest sequence durably fsynced to WAL. +func (ww *WalWriter) GetDurableSequence() uint64 { + return ww.seqManager.Durable() +} + +// IsWriteStopped reports whether the writer is rejecting new writes. +func (ww *WalWriter) IsWriteStopped() bool { + return ww.writeStopped.Load() +} + +// Close drains queued writes, stops the writer goroutine, syncs, and closes the segment. +func (ww *WalWriter) Close() error { + ww.closeOnce.Do(func() { + ww.writeStopped.Store(true) + ww.queue.Close() + close(ww.done) + ww.wg.Wait() + if err := ww.segManager.Sync(); err != nil { + ww.closeErr = err + return + } + ww.closeErr = ww.segManager.Close() + }) + return ww.closeErr +} + +func cloneBytes(src []byte) []byte { + if src == nil { + return nil + } + dst := make([]byte, len(src)) + copy(dst, src) + return dst +} diff --git a/wal/writer_test.go b/wal/writer_test.go new file mode 100644 index 0000000..25b758a --- /dev/null +++ b/wal/writer_test.go @@ -0,0 +1,120 @@ +package wal + +import ( + "fmt" + "sync" + "testing" + + "github.com/dailz/go-kv/config" +) + +func newTestWalWriter(t *testing.T) *WalWriter { + t.Helper() + + cfg := config.Defaults() + ww, err := NewWalWriter(&cfg, t.TempDir(), 0, 0) + if err != nil { + t.Fatalf("NewWalWriter: %v", err) + } + t.Cleanup(func() { + if err := ww.Close(); err != nil { + t.Fatalf("Close: %v", err) + } + }) + + return ww +} + +func TestWalWriterSinglePut(t *testing.T) { + ww := newTestWalWriter(t) + + if err := ww.Put([]byte("k1"), []byte("v1")); err != nil { + t.Fatalf("Put: %v", err) + } + + got := ww.Get([]byte("k1")) + if !got.Found { + t.Fatal("Get(k1) not found") + } + if string(got.Value) != "v1" { + t.Fatalf("Get(k1) value = %q, want %q", got.Value, "v1") + } +} + +func TestWalWriterMultiplePuts(t *testing.T) { + ww := newTestWalWriter(t) + + for i := range 10 { + key := fmt.Appendf(nil, "k%d", i) + value := fmt.Appendf(nil, "v%d", i) + if err := ww.Put(key, value); err != nil { + t.Fatalf("Put(%q): %v", key, err) + } + } + + for i := range 10 { + key := fmt.Appendf(nil, "k%d", i) + want := fmt.Sprintf("v%d", i) + got := ww.Get(key) + if !got.Found { + t.Fatalf("Get(%q) not found", key) + } + if string(got.Value) != want { + t.Fatalf("Get(%q) value = %q, want %q", key, got.Value, want) + } + } +} + +func TestWalWriterGroupCommit(t *testing.T) { + ww := newTestWalWriter(t) + + const writers = 5 + var wg sync.WaitGroup + errCh := make(chan error, writers) + + for i := range writers { + wg.Go(func() { + key := fmt.Appendf(nil, "group-k%d", i) + value := fmt.Appendf(nil, "group-v%d", i) + if err := ww.Put(key, value); err != nil { + errCh <- fmt.Errorf("put %d: %w", i, err) + } + }) + } + + wg.Wait() + close(errCh) + for err := range errCh { + if err != nil { + t.Fatal(err) + } + } + + for i := range writers { + key := fmt.Appendf(nil, "group-k%d", i) + want := fmt.Sprintf("group-v%d", i) + got := ww.Get(key) + if !got.Found { + t.Fatalf("Get(%q) not found", key) + } + if string(got.Value) != want { + t.Fatalf("Get(%q) value = %q, want %q", key, got.Value, want) + } + } +} + +func TestWalWriterDelete(t *testing.T) { + ww := newTestWalWriter(t) + + if err := ww.Put([]byte("k"), []byte("v")); err != nil { + t.Fatalf("Put: %v", err) + } + if err := ww.Delete([]byte("k")); err != nil { + t.Fatalf("Delete: %v", err) + } + + got := ww.Get([]byte("k")) + if got.Found { + t.Fatalf("Get(k) found deleted key with value %q", got.Value) + } +}