Files
go-kv/wal/writer.go
T
dailz b833a21848 fix: address Final Verification Wave findings
- Sync() now flushes BlockWriter before fd.Sync() for durability
- Pre-validate encoding before sequence allocation (design doc compliance)
- Remove dead _ = rec assignment in recover.go
- Remove unused maxPayload field from SegmentWriter
- Handle Put error in MemTable.Publish with panic on invariant violation
2026-06-12 14:29:53 +08:00

349 lines
8.2 KiB
Go

package wal
import (
"errors"
"sync"
"sync/atomic"
"time"
"github.com/dailz/go-kv/errkit"
"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, errkit.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
}
// Pre-validate encoding can succeed before sequence allocation.
if _, err := EncodeWalBatch(0, 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, errkit.ErrCommitUnknown)
return
}
active := ww.memTables.GetActive()
if err := writePending(active, entries, baseSequence); err != nil {
abortRange(active, baseSequence, lastSequence)
ww.stopWithError(requests, errkit.ErrCommitUnknown)
return
}
if err := ww.segManager.Sync(); err != nil {
abortRange(active, baseSequence, lastSequence)
ww.stopWithError(requests, errkit.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 errkit.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 errkit.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
}