diff --git a/cmd/main.go b/cmd/main.go index 4655361..a51e91b 100644 --- a/cmd/main.go +++ b/cmd/main.go @@ -23,6 +23,7 @@ func main() { readBatchSize := c.Int("read-batch-size") debug := c.Bool("debug") asyncLog := c.Bool("async-log") + storageType := c.String("storage") r := raft.New(raft.Config{ ID: id, ConfPath: conf, @@ -30,6 +31,7 @@ func main() { ReadBatchSize: readBatchSize, Debug: debug, AsyncLog: asyncLog, + StorageType: storageType, }, raft.NewKVStore()) r.Run() return nil @@ -65,6 +67,11 @@ func main() { Usage: "Enable asynchronous disk writes", Value: false, }, + &cli.StringFlag{ + Name: "storage", + Usage: "Storage backend: \"file\", \"bitcask\", or \"iouring\"", + Value: "file", + }, }, }, { diff --git a/iouring_ring.go b/iouring_ring.go new file mode 100644 index 0000000..28c55ec --- /dev/null +++ b/iouring_ring.go @@ -0,0 +1,321 @@ +package raft + +// Minimal io_uring wrapper using raw syscalls. +// +// Architecture: +// SQ ring (submission queue): +// array[i] = i (initialized once, natural mapping) +// tail is advanced by user when adding SQEs +// head is advanced by kernel when consuming SQEs +// SQE array (at IORING_OFF_SQES): +// 64-byte entries, indexed by sq.array values +// CQ ring (completion queue, shares mmap with SQ if SINGLE_MMAP): +// cqes[i] = 16-byte completion entries +// head is advanced by user after reading CQEs +// tail is advanced by kernel when posting completions + +import ( + "fmt" + "sync/atomic" + "syscall" + "unsafe" +) + +const ( + sysIoUringSetup uintptr = 425 + sysIoUringEnter uintptr = 426 + + iouringOpFsync = 3 // IORING_OP_FSYNC + iouringOpWrite = 23 // IORING_OP_WRITE (scatter/gather with offset) + + iouringFsyncDatasync = 1 // IORING_FSYNC_DATASYNC + iosqeIOLink = 4 // IOSQE_IO_LINK (1 << 2) + iouringEnterGetEvents = 1 // IORING_ENTER_GETEVENTS + + iouringFeatSingleMmap = 1 // IORING_FEAT_SINGLE_MMAP + + iouringOffSqRing uintptr = 0 + iouringOffCqRing uintptr = 0x8000000 + iouringOffSqes uintptr = 0x10000000 + + sqeSize uintptr = 64 + cqeSize uintptr = 16 +) + +// ioUringParams mirrors struct io_uring_params (120 bytes). +type ioUringParams struct { + sqEntries uint32 + cqEntries uint32 + flags uint32 + sqThreadCPU uint32 + sqThreadIdle uint32 + features uint32 + wqFD uint32 + resv [3]uint32 + sqOff ioSqRingOffsets // 40 bytes + cqOff ioCqRingOffsets // 40 bytes +} + +// ioSqRingOffsets mirrors struct io_sqring_offsets (40 bytes). +type ioSqRingOffsets struct { + head uint32 + tail uint32 + ringMask uint32 + ringEntries uint32 + flags uint32 + dropped uint32 + array uint32 + _resv1 uint32 + _userAddr uint64 +} + +// ioCqRingOffsets mirrors struct io_cqring_offsets (40 bytes). +type ioCqRingOffsets struct { + head uint32 + tail uint32 + ringMask uint32 + ringEntries uint32 + overflow uint32 + cqes uint32 + flags uint32 + _resv1 uint32 + _userAddr uint64 +} + +// ioUringSQE mirrors struct io_uring_sqe (64 bytes). +type ioUringSQE struct { + opcode uint8 + flags uint8 + ioprio uint16 + fd int32 + off uint64 // file offset for IORING_OP_WRITE + addr uint64 // buffer address + len uint32 // buffer length + rwFlags uint32 // rw_flags or fsync_flags + userData uint64 + bufIndex uint16 + personality uint16 + spliceFdIn int32 + addr3 uint64 + _pad2 uint64 +} + +// ioUringCQE mirrors struct io_uring_cqe (16 bytes). +type ioUringCQE struct { + userData uint64 + res int32 + flags uint32 +} + +// ioRing holds the mmap'd io_uring ring structures. +// All submit/reap operations must be called from a single goroutine. +type ioRing struct { + fd int + ringSize uint32 + + ringMmap []byte // SQ ring (and CQ if SINGLE_MMAP) + sqesMmap []byte // SQE array + cqMmap []byte // CQ ring (non-nil only when not SINGLE_MMAP) + + // Pointers into mmap'd memory (not heap, GC-safe) + sqHead *uint32 + sqTail *uint32 + sqMask *uint32 + sqArray unsafe.Pointer // []uint32 mapping ring slots → SQE indices + + sqes unsafe.Pointer // []ioUringSQE + + cqHead *uint32 + cqTail *uint32 + cqMask *uint32 + cqes unsafe.Pointer // []ioUringCQE +} + +func newIoRing(entries uint32) (*ioRing, error) { + var params ioUringParams + fd, _, errno := syscall.Syscall( + sysIoUringSetup, + uintptr(entries), + uintptr(unsafe.Pointer(¶ms)), + 0, + ) + if errno != 0 { + return nil, fmt.Errorf("io_uring_setup: %w", errno) + } + ringFd := int(fd) + + r := &ioRing{fd: ringFd, ringSize: params.sqEntries} + + // Compute mmap sizes. + sqRingSize := uintptr(params.sqOff.array) + uintptr(params.sqEntries)*4 + cqRingSize := uintptr(params.cqOff.cqes) + uintptr(params.cqEntries)*cqeSize + ringMmapSize := sqRingSize + if cqRingSize > ringMmapSize { + ringMmapSize = cqRingSize + } + sqesMmapSize := uintptr(params.sqEntries) * sqeSize + + // mmap SQ ring (contains CQ ring data too if SINGLE_MMAP). + ringData, err := syscall.Mmap( + ringFd, int64(iouringOffSqRing), int(ringMmapSize), + syscall.PROT_READ|syscall.PROT_WRITE, + syscall.MAP_SHARED|syscall.MAP_POPULATE, + ) + if err != nil { + syscall.Close(ringFd) + return nil, fmt.Errorf("mmap sq ring: %w", err) + } + r.ringMmap = ringData + + // mmap SQE array (always separate). + sqesData, err := syscall.Mmap( + ringFd, int64(iouringOffSqes), int(sqesMmapSize), + syscall.PROT_READ|syscall.PROT_WRITE, + syscall.MAP_SHARED|syscall.MAP_POPULATE, + ) + if err != nil { + syscall.Munmap(ringData) + syscall.Close(ringFd) + return nil, fmt.Errorf("mmap sqes: %w", err) + } + r.sqesMmap = sqesData + + // mmap CQ ring separately if not SINGLE_MMAP. + if params.features&iouringFeatSingleMmap == 0 { + cqData, err := syscall.Mmap( + ringFd, int64(iouringOffCqRing), int(cqRingSize), + syscall.PROT_READ|syscall.PROT_WRITE, + syscall.MAP_SHARED|syscall.MAP_POPULATE, + ) + if err != nil { + syscall.Munmap(sqesData) + syscall.Munmap(ringData) + syscall.Close(ringFd) + return nil, fmt.Errorf("mmap cq ring: %w", err) + } + r.cqMmap = cqData + } + + // Wire up SQ ring pointers. + sqBase := unsafe.Pointer(&ringData[0]) + r.sqHead = (*uint32)(unsafe.Pointer(uintptr(sqBase) + uintptr(params.sqOff.head))) + r.sqTail = (*uint32)(unsafe.Pointer(uintptr(sqBase) + uintptr(params.sqOff.tail))) + r.sqMask = (*uint32)(unsafe.Pointer(uintptr(sqBase) + uintptr(params.sqOff.ringMask))) + r.sqArray = unsafe.Pointer(uintptr(sqBase) + uintptr(params.sqOff.array)) + + // Wire up SQE array pointer. + r.sqes = unsafe.Pointer(&sqesData[0]) + + // Initialize sq.array[i] = i (natural 1-to-1 mapping of ring slots to SQEs). + for i := uint32(0); i < params.sqEntries; i++ { + *(*uint32)(unsafe.Pointer(uintptr(r.sqArray) + uintptr(i*4))) = i + } + + // Wire up CQ ring pointers. + var cqBase unsafe.Pointer + if r.cqMmap != nil { + cqBase = unsafe.Pointer(&r.cqMmap[0]) + } else { + cqBase = sqBase + } + r.cqHead = (*uint32)(unsafe.Pointer(uintptr(cqBase) + uintptr(params.cqOff.head))) + r.cqTail = (*uint32)(unsafe.Pointer(uintptr(cqBase) + uintptr(params.cqOff.tail))) + r.cqMask = (*uint32)(unsafe.Pointer(uintptr(cqBase) + uintptr(params.cqOff.ringMask))) + r.cqes = unsafe.Pointer(uintptr(cqBase) + uintptr(params.cqOff.cqes)) + + return r, nil +} + +// submitWriteSync submits Pwrite + Fdatasync as a linked pair and waits +// for both to complete. buf must remain valid until this returns. +func (r *ioRing) submitWriteSync(fd int, buf []byte, offset int64) error { + tail := atomic.LoadUint32(r.sqTail) + mask := atomic.LoadUint32(r.sqMask) + + // SQE 0: Write, linked to SQE 1. + sqe0 := r.sqeAt(tail & mask) + *sqe0 = ioUringSQE{} + sqe0.opcode = iouringOpWrite + sqe0.flags = iosqeIOLink + sqe0.fd = int32(fd) + sqe0.off = uint64(offset) + sqe0.addr = uint64(uintptr(unsafe.Pointer(&buf[0]))) + sqe0.len = uint32(len(buf)) + sqe0.userData = 1 + + // SQE 1: Fdatasync (tail of the linked chain, no LINK flag). + sqe1 := r.sqeAt((tail + 1) & mask) + *sqe1 = ioUringSQE{} + sqe1.opcode = iouringOpFsync + sqe1.fd = int32(fd) + sqe1.rwFlags = iouringFsyncDatasync + sqe1.userData = 2 + + atomic.StoreUint32(r.sqTail, tail+2) + return r.enterAndReap(2, 2) +} + +// submitWriteAsync submits only Pwrite (no fsync) and waits for completion. +// buf must remain valid until this returns. +func (r *ioRing) submitWriteAsync(fd int, buf []byte, offset int64) error { + tail := atomic.LoadUint32(r.sqTail) + mask := atomic.LoadUint32(r.sqMask) + + sqe := r.sqeAt(tail & mask) + *sqe = ioUringSQE{} + sqe.opcode = iouringOpWrite + sqe.fd = int32(fd) + sqe.off = uint64(offset) + sqe.addr = uint64(uintptr(unsafe.Pointer(&buf[0]))) + sqe.len = uint32(len(buf)) + sqe.userData = 1 + + atomic.StoreUint32(r.sqTail, tail+1) + return r.enterAndReap(1, 1) +} + +func (r *ioRing) sqeAt(index uint32) *ioUringSQE { + return (*ioUringSQE)(unsafe.Pointer(uintptr(r.sqes) + uintptr(index)*sqeSize)) +} + +func (r *ioRing) cqeAt(index uint32) *ioUringCQE { + return (*ioUringCQE)(unsafe.Pointer(uintptr(r.cqes) + uintptr(index)*cqeSize)) +} + +// enterAndReap submits toSubmit SQEs, waits for toWait completions, then +// reaps and checks results. +func (r *ioRing) enterAndReap(toSubmit, toWait uint32) error { + _, _, errno := syscall.Syscall6( + sysIoUringEnter, + uintptr(r.fd), + uintptr(toSubmit), + uintptr(toWait), + uintptr(iouringEnterGetEvents), + 0, 0, + ) + if errno != 0 { + return fmt.Errorf("io_uring_enter: %w", errno) + } + + head := atomic.LoadUint32(r.cqHead) + mask := atomic.LoadUint32(r.cqMask) + + var firstErr error + for i := uint32(0); i < toWait; i++ { + cqe := r.cqeAt((head + i) & mask) + if cqe.res < 0 && firstErr == nil { + firstErr = syscall.Errno(-cqe.res) + } + } + atomic.StoreUint32(r.cqHead, head+toWait) + return firstErr +} + +func (r *ioRing) close() { + if r.cqMmap != nil { + syscall.Munmap(r.cqMmap) + } + syscall.Munmap(r.sqesMmap) + syscall.Munmap(r.ringMmap) + syscall.Close(r.fd) +} diff --git a/raft.go b/raft.go index 4f57670..4f65b95 100644 --- a/raft.go +++ b/raft.go @@ -19,6 +19,7 @@ type Config struct { ReadBatchSize int // default: 128 Debug bool AsyncLog bool + StorageType string // "file" (default) or "bitcask" } type LogEntry struct { @@ -50,7 +51,7 @@ type Raft struct { pendingResponses map[int]chan Response mu sync.RWMutex peerIPPort map[int]string - storage *Storage + storage StorageBackend commitCond *sync.Cond replicating map[int]bool newLogEntryCh chan bool @@ -71,7 +72,7 @@ func New(cfg Config, sm StateMachine) *Raft { } peerIPPort := ParseConfig(cfg.ConfPath) - storage, err := NewStorage(cfg.ID, cfg.AsyncLog) + storage, err := NewStorageBackend(cfg.ID, cfg.AsyncLog, cfg.StorageType) if err != nil { panic(err) } diff --git a/storage.go b/storage.go index b30ee4c..fdc7f2c 100644 --- a/storage.go +++ b/storage.go @@ -8,7 +8,28 @@ import ( "os" ) -type Storage struct { +type StorageBackend interface { + SaveState(term int, votedFor int) error + LoadState() (int, int, error) + AppendEntry(entry LogEntry) error + AppendEntries(entries []LogEntry) error + TruncateLog(index int) error + LoadLog() ([]LogEntry, error) + Close() error +} + +func NewStorageBackend(id int, async bool, storageType string) (StorageBackend, error) { + switch storageType { + case "bitcask": + return NewBitcaskStorage(id, async) + case "iouring": + return NewIoUringStorage(id, async) + default: + return NewFileStorage(id, async) + } +} + +type FileStorage struct { id int stateFile *os.File logFile *os.File @@ -17,7 +38,7 @@ type Storage struct { async bool } -func NewStorage(id int, async bool) (*Storage, error) { +func NewFileStorage(id int, async bool) (*FileStorage, error) { stateFilename := fmt.Sprintf("raft_state_%d.bin", id) logFilename := fmt.Sprintf("raft_log_%d.bin", id) @@ -32,7 +53,7 @@ func NewStorage(id int, async bool) (*Storage, error) { return nil, err } - return &Storage{ + return &FileStorage{ id: id, stateFile: sFile, logFile: lFile, @@ -42,7 +63,7 @@ func NewStorage(id int, async bool) (*Storage, error) { }, nil } -func (s *Storage) SaveState(term int, votedFor int) error { +func (s *FileStorage) SaveState(term int, votedFor int) error { if _, err := s.stateFile.Seek(0, 0); err != nil { return err } @@ -61,7 +82,7 @@ func (s *Storage) SaveState(term int, votedFor int) error { return nil } -func (s *Storage) LoadState() (int, int, error) { +func (s *FileStorage) LoadState() (int, int, error) { info, err := s.stateFile.Stat() if err != nil { return 0, -2, err @@ -85,7 +106,7 @@ func (s *Storage) LoadState() (int, int, error) { return term, votedFor, nil } -func (s *Storage) AppendEntry(entry LogEntry) error { +func (s *FileStorage) AppendEntry(entry LogEntry) error { offset, err := s.logFile.Seek(0, io.SeekEnd) if err != nil { return err @@ -113,7 +134,7 @@ func (s *Storage) AppendEntry(entry LogEntry) error { return nil } -func (s *Storage) AppendEntries(entries []LogEntry) error { +func (s *FileStorage) AppendEntries(entries []LogEntry) error { if err := s.logWriter.Flush(); err != nil { return err } @@ -142,14 +163,13 @@ func (s *Storage) AppendEntries(entries []LogEntry) error { if err := s.logWriter.Flush(); err != nil { return err } - //return nil if !s.async { return s.logFile.Sync() } return nil } -func (s *Storage) TruncateLog(index int) error { +func (s *FileStorage) TruncateLog(index int) error { if index < 0 { return nil } @@ -181,7 +201,7 @@ func (s *Storage) TruncateLog(index int) error { return nil } -func (s *Storage) LoadLog() ([]LogEntry, error) { +func (s *FileStorage) LoadLog() ([]LogEntry, error) { if _, err := s.logFile.Seek(0, 0); err != nil { return nil, err } @@ -228,7 +248,7 @@ func (s *Storage) LoadLog() ([]LogEntry, error) { return logs, nil } -func (s *Storage) Close() error { +func (s *FileStorage) Close() error { s.logWriter.Flush() s.stateFile.Close() return s.logFile.Close() diff --git a/storage_bitcask.go b/storage_bitcask.go new file mode 100644 index 0000000..4056952 --- /dev/null +++ b/storage_bitcask.go @@ -0,0 +1,154 @@ +package raft + +import ( + "encoding/binary" + "fmt" + "io" + "sort" + + "github.com/octu0/bitcaskdb" +) + +type BitcaskStorage struct { + db *bitcaskdb.Bitcask + logCount int + async bool +} + +func NewBitcaskStorage(id int, async bool) (*BitcaskStorage, error) { + path := fmt.Sprintf("raft_bitcask_%d", id) + db, err := bitcaskdb.Open(path) + if err != nil { + return nil, err + } + return &BitcaskStorage{ + db: db, + async: async, + }, nil +} + +var stateKey = []byte("raft:state") + +const logKeyPrefix = "raft:log:" + +func logKey(index int) []byte { + return []byte(fmt.Sprintf("raft:log:%010d", index)) +} + +func (s *BitcaskStorage) SaveState(term int, votedFor int) error { + buf := make([]byte, 16) + binary.LittleEndian.PutUint64(buf[0:8], uint64(term)) + binary.LittleEndian.PutUint64(buf[8:16], uint64(votedFor)) + if err := s.db.PutBytes(stateKey, buf); err != nil { + return err + } + if !s.async { + return s.db.Sync() + } + return nil +} + +func (s *BitcaskStorage) LoadState() (int, int, error) { + if !s.db.Has(stateKey) { + return 0, -2, nil + } + rc, err := s.db.Get(stateKey) + if err != nil { + return 0, -2, err + } + defer rc.Close() + buf, err := io.ReadAll(rc) + if err != nil { + return 0, 0, err + } + if len(buf) < 16 { + return 0, -2, nil + } + term := int(binary.LittleEndian.Uint64(buf[0:8])) + votedFor := int(binary.LittleEndian.Uint64(buf[8:16])) + return term, votedFor, nil +} + +func (s *BitcaskStorage) appendEntryRaw(entry LogEntry) error { + key := logKey(s.logCount) + buf := make([]byte, 8+len(entry.Command)) + binary.LittleEndian.PutUint64(buf[0:8], uint64(entry.Term)) + copy(buf[8:], entry.Command) + if err := s.db.PutBytes(key, buf); err != nil { + return err + } + s.logCount++ + return nil +} + +func (s *BitcaskStorage) AppendEntry(entry LogEntry) error { + if err := s.appendEntryRaw(entry); err != nil { + return err + } + if !s.async { + return s.db.Sync() + } + return nil +} + +func (s *BitcaskStorage) AppendEntries(entries []LogEntry) error { + for _, entry := range entries { + if err := s.appendEntryRaw(entry); err != nil { + return err + } + } + if !s.async { + return s.db.Sync() + } + return nil +} + +func (s *BitcaskStorage) TruncateLog(index int) error { + for i := index; i < s.logCount; i++ { + if err := s.db.Delete(logKey(i)); err != nil { + return err + } + } + s.logCount = index + if !s.async { + return s.db.Sync() + } + return nil +} + +func (s *BitcaskStorage) LoadLog() ([]LogEntry, error) { + var keys []string + err := s.db.Scan([]byte(logKeyPrefix), func(key []byte) error { + keys = append(keys, string(key)) + return nil + }) + if err != nil { + return nil, err + } + sort.Strings(keys) + + var logs []LogEntry + for _, k := range keys { + rc, err := s.db.Get([]byte(k)) + if err != nil { + return nil, err + } + data, err := io.ReadAll(rc) + rc.Close() + if err != nil { + return nil, err + } + if len(data) < 8 { + return nil, fmt.Errorf("corrupted log entry: %s", k) + } + term := int(binary.LittleEndian.Uint64(data[0:8])) + cmd := data[8:] + logs = append(logs, LogEntry{Term: term, Command: cmd}) + } + s.logCount = len(logs) + return logs, nil +} + +func (s *BitcaskStorage) Close() error { + return s.db.Close() +} diff --git a/storage_iouring.go b/storage_iouring.go new file mode 100644 index 0000000..ca15a45 --- /dev/null +++ b/storage_iouring.go @@ -0,0 +1,348 @@ +package raft + +// IoUringStorage is a StorageBackend that uses io_uring for log writes. +// +// Architecture: +// +// AppendEntry / AppendEntries +// │ send writeJob to writeCh +// ↓ +// submitter goroutine (owns ioRing) +// │ drain all pending writeJobs (group commit) +// │ concatenate data → one Pwrite + one Fdatasync via io_uring +// ↓ +// notify each caller with its starting file offset +// +// Group commit happens naturally: while one Pwrite+Fdatasync is in flight +// in io_uring_enter, new jobs accumulate in writeCh. On the next loop +// iteration the submitter drains all of them and issues a single fsync +// covering every entry. + +import ( + "bufio" + "encoding/binary" + "fmt" + "io" + "os" + "sync" +) + +// writeJob is sent by AppendEntry / AppendEntries to the submitter goroutine. +type writeJob struct { + data []byte // pre-encoded log bytes (one or more entries) + resultCh chan writeResult +} + +// writeResult is returned by the submitter to the caller. +type writeResult struct { + startOffset int64 // file offset where this job's data begins + err error +} + +// truncateJob is sent by TruncateLog to the submitter goroutine. +type truncateJob struct { + at int64 + doneCh chan error +} + +// IoUringStorage implements StorageBackend using io_uring. +type IoUringStorage struct { + stateFile *os.File + logFile *os.File + logFd int + logOffsets []int64 // maintained by AppendEntry/AppendEntries callers + async bool + + ring *ioRing + writeCh chan writeJob + truncateCh chan truncateJob + stopCh chan struct{} + wg sync.WaitGroup +} + +func NewIoUringStorage(id int, async bool) (*IoUringStorage, error) { + stateFilename := fmt.Sprintf("raft_state_%d.bin", id) + logFilename := fmt.Sprintf("raft_log_%d.bin", id) + + sFile, err := os.OpenFile(stateFilename, os.O_RDWR|os.O_CREATE, 0644) + if err != nil { + return nil, err + } + + // O_RDWR|O_CREAT for log; O_DIRECT is omitted to keep aligned-write + // complexity out of scope. + lFile, err := os.OpenFile(logFilename, os.O_RDWR|os.O_CREATE, 0644) + if err != nil { + sFile.Close() + return nil, err + } + + info, err := lFile.Stat() + if err != nil { + sFile.Close() + lFile.Close() + return nil, err + } + + ring, err := newIoRing(16) + if err != nil { + sFile.Close() + lFile.Close() + return nil, fmt.Errorf("io_uring unavailable: %w", err) + } + + s := &IoUringStorage{ + stateFile: sFile, + logFile: lFile, + logFd: int(lFile.Fd()), + logOffsets: []int64{}, + async: async, + ring: ring, + writeCh: make(chan writeJob, 256), + truncateCh: make(chan truncateJob, 1), + stopCh: make(chan struct{}), + } + + s.wg.Add(1) + go s.runSubmitter(info.Size()) + return s, nil +} + +// runSubmitter is the single goroutine that owns the ioRing. +func (s *IoUringStorage) runSubmitter(initialOffset int64) { + defer s.wg.Done() + currentOffset := initialOffset + + for { + // Wait for the first event. + var jobs []writeJob + select { + case job := <-s.writeCh: + jobs = append(jobs, job) + case trunc := <-s.truncateCh: + err := s.doTruncate(trunc.at) + if err == nil { + currentOffset = trunc.at + } + trunc.doneCh <- err + continue + case <-s.stopCh: + return + } + + // Drain any additional write jobs that are already queued + // (group commit: batch them into one io_uring submission). + drain: + for { + select { + case job := <-s.writeCh: + jobs = append(jobs, job) + default: + break drain + } + } + + // Calculate per-job starting offsets and concatenate all data. + offsets := make([]int64, len(jobs)) + off := currentOffset + var totalSize int + for _, j := range jobs { + totalSize += len(j.data) + } + allData := make([]byte, 0, totalSize) + for i, j := range jobs { + offsets[i] = off + allData = append(allData, j.data...) + off += int64(len(j.data)) + } + + // Submit to io_uring. + var err error + if s.async { + err = s.ring.submitWriteAsync(s.logFd, allData, currentOffset) + } else { + err = s.ring.submitWriteSync(s.logFd, allData, currentOffset) + } + + if err == nil { + currentOffset = off + } + + // Notify all callers with their individual starting offsets. + for i, j := range jobs { + startOff := offsets[i] + if err != nil { + startOff = -1 + } + j.resultCh <- writeResult{startOffset: startOff, err: err} + } + } +} + +func (s *IoUringStorage) doTruncate(at int64) error { + if err := s.logFile.Truncate(at); err != nil { + return err + } + if !s.async { + return s.logFile.Sync() + } + return nil +} + +// SaveState uses regular file I/O (state writes are infrequent). +func (s *IoUringStorage) SaveState(term int, votedFor int) error { + if _, err := s.stateFile.Seek(0, 0); err != nil { + return err + } + buf := make([]byte, 16) + binary.LittleEndian.PutUint64(buf[0:8], uint64(term)) + binary.LittleEndian.PutUint64(buf[8:16], uint64(votedFor)) + if _, err := s.stateFile.Write(buf); err != nil { + return err + } + if !s.async { + return s.stateFile.Sync() + } + return nil +} + +func (s *IoUringStorage) LoadState() (int, int, error) { + info, err := s.stateFile.Stat() + if err != nil { + return 0, -2, err + } + if info.Size() == 0 { + return 0, -2, nil + } + if _, err := s.stateFile.Seek(0, 0); err != nil { + return 0, 0, err + } + buf := make([]byte, 16) + if _, err := io.ReadFull(s.stateFile, buf); err != nil { + return 0, 0, err + } + term := int(binary.LittleEndian.Uint64(buf[0:8])) + votedFor := int(binary.LittleEndian.Uint64(buf[8:16])) + return term, votedFor, nil +} + +// encodeLogEntry serialises a LogEntry to the same binary format as +// FileStorage: term (int64 LE) + cmdLen (int64 LE) + command bytes. +func encodeLogEntry(entry LogEntry) []byte { + buf := make([]byte, 16+len(entry.Command)) + binary.LittleEndian.PutUint64(buf[0:8], uint64(entry.Term)) + binary.LittleEndian.PutUint64(buf[8:16], uint64(len(entry.Command))) + copy(buf[16:], entry.Command) + return buf +} + +func (s *IoUringStorage) AppendEntry(entry LogEntry) error { + data := encodeLogEntry(entry) + resultCh := make(chan writeResult, 1) + s.writeCh <- writeJob{data: data, resultCh: resultCh} + res := <-resultCh + if res.err != nil { + return res.err + } + s.logOffsets = append(s.logOffsets, res.startOffset) + return nil +} + +func (s *IoUringStorage) AppendEntries(entries []LogEntry) error { + if len(entries) == 0 { + return nil + } + + // Encode all entries into one contiguous buffer so the submitter can + // write them in a single Pwrite. + entrySizes := make([]int, len(entries)) + var allData []byte + for i, e := range entries { + enc := encodeLogEntry(e) + entrySizes[i] = len(enc) + allData = append(allData, enc...) + } + + resultCh := make(chan writeResult, 1) + s.writeCh <- writeJob{data: allData, resultCh: resultCh} + res := <-resultCh + if res.err != nil { + return res.err + } + + // Track the starting offset of each individual entry. + off := res.startOffset + for i := range entries { + s.logOffsets = append(s.logOffsets, off) + off += int64(entrySizes[i]) + } + return nil +} + +func (s *IoUringStorage) TruncateLog(index int) error { + if index < 0 { + return nil + } + if index >= len(s.logOffsets) { + return nil + } + truncateAt := s.logOffsets[index] + doneCh := make(chan error, 1) + s.truncateCh <- truncateJob{at: truncateAt, doneCh: doneCh} + if err := <-doneCh; err != nil { + return err + } + s.logOffsets = s.logOffsets[:index] + return nil +} + +// LoadLog reads the log file sequentially using buffered I/O. +// Called only at startup, before any writes begin. +func (s *IoUringStorage) LoadLog() ([]LogEntry, error) { + if _, err := s.logFile.Seek(0, 0); err != nil { + return nil, err + } + + var logs []LogEntry + s.logOffsets = []int64{} + + reader := bufio.NewReader(s.logFile) + offset := int64(0) + + for { + startOffset := offset + + var term int64 + if err := binary.Read(reader, binary.LittleEndian, &term); err == io.EOF { + break + } else if err != nil { + return nil, err + } + offset += 8 + + var cmdLen int64 + if err := binary.Read(reader, binary.LittleEndian, &cmdLen); err != nil { + return nil, err + } + offset += 8 + + cmd := make([]byte, cmdLen) + if _, err := io.ReadFull(reader, cmd); err != nil { + return nil, err + } + offset += cmdLen + + s.logOffsets = append(s.logOffsets, startOffset) + logs = append(logs, LogEntry{Term: int(term), Command: cmd}) + } + + return logs, nil +} + +func (s *IoUringStorage) Close() error { + close(s.stopCh) + s.wg.Wait() + s.ring.close() + s.stateFile.Close() + return s.logFile.Close() +}