package wal import ( "errors" "sync" "sync/atomic" "time" "github.com/dailz/go-kv/config" "github.com/dailz/go-kv/errkit" "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() const batchSizeThreshold = 32 * 1024 for { select { case <-ww.done: ww.processRemaining() return case req, ok := <-ww.queue.ch: if !ok { ww.processRemaining() return } requests := []*CommitRequest{req} accumulatedSize := estimateRequestSize(req) timer := time.NewTimer(ww.cfg.GroupCommitDelay) collect: for accumulatedSize < batchSizeThreshold { select { case <-timer.C: break collect case req, ok := <-ww.queue.ch: if !ok { break collect } requests = append(requests, req) accumulatedSize += estimateRequestSize(req) case <-ww.done: if !timer.Stop() { select { case <-timer.C: default: } } ww.sendError(requests, errkit.ErrWriteStopped) ww.processRemaining() return } } if !timer.Stop() { select { case <-timer.C: default: } } ww.processBatch(requests) } } } func estimateRequestSize(req *CommitRequest) int { size := 0 for _, entry := range req.Entries { size += len(entry.Key) + len(entry.Value) + 4 } return size } 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, baseSequence); 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 }