feat(wal): implement WAL writer with group commit and recovery batch replay

- wal/commit_queue.go: bounded buffered channel for write requests
- wal/writer.go: single-goroutine main loop implementing 11-step write flow
  with group commit, sequence allocation, MemTable publish/abort, write-stopped
- wal/recovery.go: BatchReplayer interface, ReplayBatch, ReplaySegmentFile,
  RecoverFromSegments with fragment reassembly and tail corruption handling
- Comprehensive tests for all modules, all pass with -race
This commit is contained in:
dailz
2026-06-12 13:57:31 +08:00
parent 08960a9bcf
commit fe2d4fc5f0
7 changed files with 971 additions and 4 deletions
+28 -2
View File
@@ -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"
+2 -2
View File
@@ -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`:恢复核心逻辑
+58
View File
@@ -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)
}
+169
View File
@@ -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
}
+252
View File
@@ -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)
}
}
+342
View File
@@ -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
}
+120
View File
@@ -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)
}
}