-
Notifications
You must be signed in to change notification settings - Fork 34
Raft集成设计文档
本文档描述如何将 openraft 集成到 kiwi-rs 项目中,实现与 C++ 版本兼容的 Raft 协议支持。
- 操作级日志:使用 Binlog(操作级)而非命令级日志,与 C++ 版本保持一致
- 每个 CF 一个 Index:维护每个 ColumnFamily 的 applied_index 和 flushed_index
- Batch 抽象层:统一 Raft 模式和单机模式的接口
- Leader 检查:读写命令都需要在 Leader 执行
客户端请求
│
├─→ NetworkHandler
│ │
│ └─→ CmdTable::execute()
│ │
│ ├─→ 【Leader 检查】
│ │
│ └─→ Cmd::do_cmd()
│ │
│ └─→ Storage::set/get()
│ │
│ ├─→ 写操作:创建 Batch
│ │ ├─→ RaftBatch → RaftNode::append_log()
│ │ └─→ RocksBatch → 直接写入
│ │
│ └─→ 读操作:直接读取 RocksDB
│
└─→ Raft 层
│
├─→ LogStore (RaftLogStorage)
│ └─→ 存储 Raft 日志到文件系统(与 C++ 版本一致)
│
├─→ StateMachine (RaftStateMachine)
│ └─→ 应用 Binlog 到 Storage
│
└─→ FlushEventListener
└─→ 监听 Flush 事件,触发快照
src/
├── raft/
│ ├── mod.rs # Raft 模块入口
│ ├── types.rs # TypeConfig 定义
│ ├── log_store.rs # RaftLogStorage + RaftLogReader(文件系统)
│ ├── state_machine.rs # RaftStateMachine + RaftSnapshotBuilder
│ ├── log_index.rs # LogIndex 管理(每个 CF 一个 index)
│ ├── flush_manager.rs # Flush 事件监听器
│ ├── snapshot.rs # 快照管理
│ ├── network.rs # RaftNetworkFactory + RaftNetwork
│ ├── node.rs # RaftNode 封装
│ └── error.rs # 错误类型定义
└── storage/
└── batch.rs # Batch 抽象层
db_path/
├── 0/ # DB 0 的数据目录(RocksDB)
│ ├── 000004.log
│ ├── CURRENT
│ ├── MANIFEST-000005
│ └── ...
└── 0/_praft/ # DB 0 的 Raft 目录(文件系统)
├── log/ # Raft 日志文件(段模式)
│ ├── 0000000000000001.log # 段文件:包含 log_index 1 到 N 的 Entry
│ ├── 0000000000000100.log # 段文件:包含 log_index 100 到 M 的 Entry
│ ├── segments.json # 段元数据:记录每个段的起始和结束 log_index
│ └── ...
└── raft_meta/ # Raft 元数据文件
├── vote # Vote 信息
├── committed # Committed log ID
└── last_purged_log_id # 最后清理的日志 ID
日志文件组织说明:
- 段(Segment)模式:一个文件包含多个连续的 log_index 的 Entry
-
文件命名:基于段的起始 log_index:
{start_index:016x}.log - 文件大小限制:每个段文件最大 64MB(可配置),达到限制时创建新段
-
段元数据:
segments.json记录每个段的起始和结束 log_index,用于快速查找
使用 openraft 的 declare_raft_types! 宏定义类型配置。
实现日志存储,将 Raft 日志持久化到文件系统(与 C++ 版本一致)。
实现状态机,应用 Binlog 到 Storage。
维护每个 CF 的 applied_index 和 flushed_index。
统一 Raft 模式和单机模式的接口。
本文档中包含多个 TODO 事项需要实现。以下是 C++ 版本、openraft 示例和推荐方案的对比:
| TODO 事项 | C++ 版本 | openraft 示例 | 推荐方案 |
|---|---|---|---|
| LogIndex 初始化 | TablePropertiesCollector | 无(内存状态) | 使用 TableProperties API |
| 快照构建 | RocksDB Checkpoint | 序列化内存状态 | 优先 Checkpoint,否则序列化 |
| 网络层 | brpc(gRPC) | HTTP/WebSocket | gRPC(tonic) |
| 状态机响应 | 通过回调 |
client_write 返回 |
使用 result.data
|
| Instance 获取 | 直接访问数组 | 无(单实例) | 存储 HashMap 映射 |
| Leader 地址 | 从配置获取 | 从配置获取 | 存储节点配置映射 |
高优先级(核心功能):
- ✅ 网络层实现(gRPC - tonic)
- ✅ 快照构建和恢复(RocksDB Checkpoint)
- ✅ 状态机响应获取(
result.data)
中优先级(性能优化): 4. LogIndex 初始化(TableProperties) 5. Batch 写入优化
低优先级(完善功能): 6. Leader 地址获取(节点配置映射) 7. Instance 获取(HashMap 映射) 8. 初始化检查
优先参考 C++ 版本的设计(TableProperties、Checkpoint、gRPC),同时参考 openraft 示例的简化实现(状态机响应、节点配置)。
文件:src/raft/src/types.rs
// src/raft/src/types.rs
use std::io::Cursor;
use bytes::Bytes;
use serde::{Deserialize, Serialize};
/// Node ID 类型
pub type NodeId = u64;
/// Node 信息(存储节点地址)
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq, Default)]
pub struct Node {
pub rpc_addr: String, // Raft RPC 地址
pub api_addr: String, // HTTP API 地址
}
/// Protobuf 生成的 Binlog 类型(与 C++ 版本一致)
/// 从 binlog.proto 生成,包含 Binlog、BinlogEntry、OperateType
pub mod binlog_proto {
include!(concat!(env!("OUT_DIR"), "/kiwi.binlog.rs"));
}
// 导出 protobuf 生成的类型
pub use binlog_proto::{Binlog, BinlogEntry, OperateType};
/// Request 类型:使用 Binlog(操作级,protobuf 类型)
/// 注意:这里 Request 就是 Binlog,不是 Redis 命令!
pub type Request = Binlog;
/// Response 类型
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct Response {
pub success: bool,
pub error: Option<String>,
}
/// 快照数据类型(使用文件句柄,支持分块传输,与 C++ 版本一致)
/// 注意:不使用 generic-snapshot-data feature,使用默认的分块传输机制
use tokio::fs::File;
pub type SnapshotData = File;
/// 使用 openraft 宏定义 TypeConfig
openraft::declare_raft_types!(
pub TypeConfig:
D = Request, // 应用数据 = Binlog(操作级,protobuf)
R = Response, // 应用响应
Node = Node, // 节点信息
NodeId = NodeId, // 节点 ID
SnapshotData = SnapshotData, // 快照数据
);
/// 类型别名
pub mod typ {
use super::*;
use openraft::error::Infallible;
pub type Entry = openraft::Entry<TypeConfig>;
pub type RaftError<E = Infallible> = openraft::error::RaftError<NodeId, E>;
pub type RPCError<E = Infallible> = openraft::error::RPCError<NodeId, Node, RaftError<E>>;
pub type ClientWriteError = openraft::error::ClientWriteError<NodeId, Node>;
pub type ClientWriteResponse = openraft::raft::ClientWriteResponse<TypeConfig>;
pub type StorageError = openraft::StorageError<NodeId>;
}Protobuf 定义文件:src/raft/proto/binlog.proto(与 C++ 版本一致)
syntax = "proto3";
package kiwi;
option optimize_for = LITE_RUNTIME;
enum OperateType {
kNoOperate = 0;
kPut = 1;
kDelete = 2;
}
message BinlogEntry {
uint32 cf_idx = 1;
OperateType op_type = 2;
bytes key = 3;
optional bytes value = 4;
}
message Binlog {
uint32 db_id = 1;
uint32 slot_idx = 2;
repeated BinlogEntry entries = 3;
}依赖配置:src/raft/Cargo.toml
[package]
name = "raft"
version.workspace = true
edition.workspace = true
[dependencies]
openraft = { workspace = true, features = ["serde", "storage-v2"] } # 不使用 generic-snapshot-data,使用默认分块传输
bytes = { workspace = true }
serde = { workspace = true, features = ["derive"] }
prost = "0.13"
prost-types = "0.13"
tonic = "0.11"
bincode = "1.3" # 仅用于 gRPC 网络层序列化 openraft 的内部 RPC 类型(AppendEntriesRequest 等)
tokio = { workspace = true, features = ["fs"] } # 用于 File 类型
rocksdb = { workspace = true }
tokio::sync = { workspace = true }
[build-dependencies]
prost-build = { version = "0.13", features = ["serde"] }
tonic-build = "0.11"构建脚本:src/raft/build.rs
fn main() -> Result<(), Box<dyn std::error::Error>> {
// 编译 binlog.proto,生成 Rust 代码
// 启用 serde 支持,以满足 openraft 的 Serialize/Deserialize 要求
prost_build::Config::new()
.type_attribute(".kiwi.Binlog", "#[derive(serde::Serialize, serde::Deserialize)]")
.type_attribute(".kiwi.BinlogEntry", "#[derive(serde::Serialize, serde::Deserialize)]")
.type_attribute(".kiwi.OperateType", "#[derive(serde::Serialize, serde::Deserialize)]")
.compile_protos(&["proto/binlog.proto"], &["proto/"])?;
// 编译 gRPC 服务定义
tonic_build::compile_protos("proto/raft.proto")?;
Ok(())
}说明:
- ✅ 与 C++ 版本一致:使用相同的
binlog.proto文件,确保跨语言兼容性 - ✅ gRPC 原生支持:protobuf 类型可直接用于 gRPC 消息,无需额外转换
- ✅ openraft 兼容:通过
prost-build的serde特性添加Serialize/Deserializetraits,满足 openraft 的序列化要求 - ✅ 无需手动转换:避免了 Rust 结构体与 protobuf 之间的转换开销,提高性能
- ✅ 类型安全:编译时生成类型,减少运行时错误
- ✅ 跨语言兼容:与 C++ 版本使用相同的数据格式,便于调试和迁移
序列化机制说明:
-
openraft 日志序列化:openraft 使用
serde_json(当启用serdefeature 时)序列化整个Entry,包括EntryPayload::Normal(Binlog) -
Binlog 序列化:由于
Binlog通过prost-build的serde特性实现了Serialize/Deserialize,openraft 可以直接序列化它 -
无需手动转换:
client_write直接接受Binlog类型,EntryPayload::Normal也直接包含Binlog类型,无需先转换为 bytes -
与 C++ 的区别:C++ 版本需要手动将
Binlog序列化为 protobuf bytes 后传给 braft,而 Rust 版本通过 serde 自动处理
为什么使用 Protobuf 而不是手动定义 Rust 结构体?
- 一致性:C++ 版本使用 protobuf,Rust 版本也应该使用相同的定义
- 性能:protobuf 序列化性能优于 serde(bincode)
- 兼容性:gRPC 原生支持 protobuf,无需额外转换
-
维护性:单一数据源(
.proto文件),避免手动维护对应关系
文件:src/storage/src/db.rs(新建)
// src/storage/src/db.rs
use std::path::PathBuf;
use std::sync::Arc;
use tokio::sync::RwLock;
use storage::Storage;
use storage::StorageOptions;
use rocksdb::checkpoint::Checkpoint;
/// DB 包装结构体(对应 C++ 的 DB 类)
/// 包含 storage_mutex 和 storage,与 C++ 版本一致
pub struct DB {
db_index: usize,
db_path: PathBuf,
/// 如果只想改变指向 storage 的指针,必须先获取 mutex 锁
/// 如果只想访问指针,只需要获取共享锁
/// (与 C++ 版本的注释一致)
storage_mutex: Arc<RwLock<()>>, // 对应 C++ 的 std::shared_mutex
/// Storage 实例(对应 C++ 的 std::unique_ptr<storage::Storage>)
storage: Arc<RwLock<Option<Arc<Storage>>>>,
opened: Arc<std::sync::atomic::AtomicBool>,
}
impl DB {
pub fn new(db_index: usize, db_path: PathBuf) -> Self {
Self {
db_index,
db_path,
storage_mutex: Arc::new(RwLock::new(())),
storage: Arc::new(RwLock::new(None)),
opened: Arc::new(std::sync::atomic::AtomicBool::new(false)),
}
}
/// 打开数据库(对应 C++ 的 DB::Open)
pub async fn open(
&self,
storage_options: Arc<StorageOptions>,
) -> Result<(), Box<dyn std::error::Error>> {
// 获取独占锁(对应 C++ 的 std::lock_guard)
let _guard = self.storage_mutex.write().await;
// 关闭旧 Storage(如果存在)
{
let mut storage_guard = self.storage.write().await;
if let Some(old_storage) = storage_guard.take() {
old_storage.shutdown().await;
drop(old_storage);
}
}
// 创建新的 Storage
let mut new_storage = Storage::new(
storage_options.db_instance_num,
storage_options.db_id,
);
// 打开 Storage
let _receiver = new_storage.open(storage_options.clone(), &self.db_path)?;
let storage_arc = Arc::new(new_storage);
// 保存 Storage
{
let mut storage_guard = self.storage.write().await;
*storage_guard = Some(storage_arc);
}
self.opened.store(true, std::sync::atomic::Ordering::SeqCst);
Ok(())
}
/// 获取 Storage(对应 C++ 的 GetStorage)
/// 使用共享锁(对应 C++ 的 LockShared/UnLockShared)
pub async fn get_storage(&self) -> Option<Arc<Storage>> {
let _guard = self.storage_mutex.read().await; // 共享锁
self.storage.read().await.clone()
}
/// 创建 Checkpoint(对应 C++ 的 CreateCheckpoint)
/// 调用 Storage 层的 create_checkpoint 方法,由 Storage 管理多个 Redis 实例的 checkpoint
pub async fn create_checkpoint(
&self,
checkpoint_path: &std::path::Path,
sync: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let checkpoint_sub_path = checkpoint_path.join(self.db_index.to_string());
std::fs::create_dir_all(&checkpoint_sub_path)?;
// 使用共享锁(对应 C++ 的 std::shared_lock)
let _guard = self.storage_mutex.read().await;
if let Some(storage) = self.storage.read().await.as_ref() {
// 调用 Storage 的 create_checkpoint 方法(对应 C++ 的 Storage::CreateCheckpoint)
// Storage 会为每个 Redis 实例创建 checkpoint
let handles = storage.create_checkpoint(&checkpoint_sub_path).await?;
// 如果 sync 为 true,等待所有任务完成(对应 C++ 的 r.get())
if sync {
for handle in handles {
handle.await??;
}
}
// 如果 sync 为 false,不等待,直接返回(对应 C++ 版本的行为)
}
Ok(())
}
/// 从 Checkpoint 加载数据库(对应 C++ 的 LoadDBFromCheckpoint)
/// 使用独占锁(对应 C++ 的 std::lock_guard)
/// 注意:C++ 版本中 sync 参数虽然标记为 [[maybe_unused]],但实际总是等待完成
pub async fn load_db_from_checkpoint(
&self,
checkpoint_path: &std::path::Path,
storage_options: Arc<StorageOptions>,
sync: bool,
) -> Result<(), Box<dyn std::error::Error>> {
let checkpoint_sub_path = checkpoint_path.join(self.db_index.to_string());
if !checkpoint_sub_path.exists() {
return Err(format!("Checkpoint dir {} does not exist!", checkpoint_sub_path.display()).into());
}
if !self.db_path.exists() {
std::fs::create_dir_all(&self.db_path)?;
}
// 获取独占锁(对应 C++ 的 std::lock_guard<std::shared_mutex>)
let _guard = self.storage_mutex.write().await;
// 标记为未打开
self.opened.store(false, std::sync::atomic::Ordering::SeqCst);
// 关闭旧 Storage(对应 C++ 的 old_storage->Close())
{
let mut storage_guard = self.storage.write().await;
if let Some(old_storage) = storage_guard.take() {
old_storage.shutdown().await;
drop(old_storage);
}
}
// 注意:C++ 版本中,备份逻辑是在 Storage::LoadCheckpointInternal 中处理的
// 每个实例的备份(重命名为 .tmp)是在复制文件之前完成的
// 所以这里不需要在 DB 层处理整个 db_path 的备份
// 创建新的 Storage(对应 C++ 的 storage_ = std::make_unique<storage::Storage>())
// 注意:此时 Storage 还未打开,insts 为空
let mut new_storage = Storage::new(
storage_options.db_instance_num,
storage_options.db_id,
);
// 调用 Storage 的 load_checkpoint 方法(对应 C++ 的 Storage::LoadCheckpoint)
// 注意:在 Storage 打开之前调用,只复制文件,不依赖 Storage 实例
// 需要传入 db_instance_num,因为此时 Storage 还没打开,insts 为空
let handles = Storage::load_checkpoint(
&checkpoint_sub_path,
&self.db_path,
storage_options.db_instance_num,
).await?;
// 等待所有加载任务完成(对应 C++ 的 r.get())
// 注意:C++ 版本中总是等待 LoadCheckpoint 完成,无论 sync 参数
for handle in handles {
handle.await??;
}
// 打开 Storage(对应 C++ 的 storage_->Open())
// 此时 checkpoint 文件已经复制完成,可以打开 Storage
let _receiver = new_storage.open(storage_options.clone(), &self.db_path)?;
let storage_arc = Arc::new(new_storage);
// 保存 Storage
{
let mut storage_guard = self.storage.write().await;
*storage_guard = Some(storage_arc);
}
// 标记为已打开
self.opened.store(true, std::sync::atomic::Ordering::SeqCst);
// 注意:C++ 版本中,备份目录的删除是在 Storage::LoadCheckpointInternal 中完成的
// 每个实例的备份目录(.tmp)在复制成功后就被删除了,不需要在 DB 层处理
Ok(())
}
pub fn get_db_index(&self) -> usize {
self.db_index
}
}
// 辅助函数:递归复制目录(对应 C++ 的 RecursiveLinkAndCopy)
// 注意:这个函数应该在 Storage 模块中实现,这里仅作为示例
fn copy_dir_all(src: impl AsRef<std::path::Path>, dst: impl AsRef<std::path::Path>) -> std::io::Result<()> {
std::fs::create_dir_all(&dst)?;
for entry in std::fs::read_dir(src)? {
let entry = entry?;
let ty = entry.file_type()?;
if ty.is_dir() {
copy_dir_all(entry.path(), dst.as_ref().join(entry.file_name()))?;
} else {
std::fs::copy(entry.path(), dst.as_ref().join(entry.file_name()))?;
}
}
Ok(())
}说明:
- ✅ 与 C++ 版本一致:
storage_mutex放在DB结构体中,与 C++ 的db.h一致 - ✅ 锁的使用:
CreateCheckpoint使用共享锁,LoadDBFromCheckpoint使用独占锁 - ✅ 函数命名:
create_checkpoint和load_db_from_checkpoint与 C++ 版本一致 - ✅ 原子操作:使用
rename确保数据一致性,支持回滚 - ✅ 分层设计:
DB层调用Storage层的 checkpoint 方法,由Storage管理多个 Redis 实例 - ✅ Checkpoint 结构:
checkpoint_path/db_index/0/,checkpoint_path/db_index/1/, ... 对应多个 RocksDB 实例
Raft checkpoint 与多个 storage 的对应关系:
Raft (一个 Raft 实例)
└─→ DB (一个 DB,对应一个 db_id)
└─→ Storage (一个 Storage)
└─→ insts (多个 Redis 实例,db_instance_num 个)
├─→ insts[0] → RocksDB 0 (路径: db_path/0/)
├─→ insts[1] → RocksDB 1 (路径: db_path/1/)
└─→ insts[N] → RocksDB N (路径: db_path/N/)
Checkpoint 目录结构:
checkpoint_path/
└─→ db_index/ # DB 的索引(对应 db_id)
├─→ 0/ # Redis 实例 0 的 checkpoint
├─→ 1/ # Redis 实例 1 的 checkpoint
└─→ N/ # Redis 实例 N 的 checkpoint
文件:src/raft/src/log_index.rs
// src/raft/src/log_index.rs
use std::collections::VecDeque;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use tokio::sync::RwLock;
use rocksdb::{DB, FlushJobInfo, EventListener};
/// LogIndex 类型
pub type LogIndex = u64;
/// SequenceNumber 类型(RocksDB)
pub type SequenceNumber = u64;
/// LogIndex 和 SequenceNumber 对(对应 C++ 的 LogIndexSeqnoPair)
#[derive(Debug, Clone)]
pub struct LogIndexSeqnoPair {
pub log_index: Arc<AtomicU64>,
pub seqno: Arc<AtomicU64>,
}
impl LogIndexSeqnoPair {
pub fn new(log_index: u64, seqno: u64) -> Self {
Self {
log_index: Arc::new(AtomicU64::new(log_index)),
seqno: Arc::new(AtomicU64::new(seqno)),
}
}
pub fn get_log_index(&self) -> u64 {
self.log_index.load(Ordering::Acquire)
}
pub fn get_seqno(&self) -> u64 {
self.seqno.load(Ordering::Acquire)
}
pub fn set(&self, log_index: u64, seqno: u64) {
self.log_index.store(log_index, Ordering::Release);
self.seqno.store(seqno, Ordering::Release);
}
pub fn max_log_index(&self, other: u64) -> u64 {
self.log_index.load(Ordering::Acquire).max(other)
}
pub fn max_seqno(&self, other: u64) -> u64 {
self.seqno.load(Ordering::Acquire).max(other)
}
}
impl PartialEq for LogIndexSeqnoPair {
fn eq(&self, other: &Self) -> bool {
self.seqno.load(Ordering::Acquire) == other.seqno.load(Ordering::Acquire)
}
}
impl PartialOrd for LogIndexSeqnoPair {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
self.seqno.load(Ordering::Acquire).partial_cmp(&other.seqno.load(Ordering::Acquire))
}
}
/// 每个 CF 的索引对(对应 C++ 的 LogIndexPair)
#[derive(Debug, Clone)]
pub struct ColumnFamilyLogIndexPair {
pub applied_index: LogIndexSeqnoPair, // memtable 中的最新记录
pub flushed_index: LogIndexSeqnoPair, // SST 文件中的最新记录
}
/// 所有 CF 的索引管理(对应 C++ 的 LogIndexOfColumnFamilies)
pub struct LogIndexOfColumnFamilies {
/// 每个 CF 的索引对
cf_indexes: Arc<RwLock<Vec<ColumnFamilyLogIndexPair>>>,
/// 全局最小 flushed_index(对应 C++ 的 last_flush_index_)
last_flush_index: Arc<RwLock<LogIndexSeqnoPair>>,
}
/// 最小索引结果(对应 C++ 的 SmallestIndexRes)
#[derive(Debug)]
pub struct SmallestIndexRes {
pub smallest_applied_log_index_cf: usize,
pub smallest_applied_log_index: LogIndex,
pub smallest_flushed_log_index_cf: usize,
pub smallest_flushed_log_index: LogIndex,
pub smallest_flushed_seqno: SequenceNumber,
}
impl LogIndexOfColumnFamilies {
pub fn new(cf_count: usize) -> Self {
let mut cf_indexes = Vec::new();
for _ in 0..cf_count {
cf_indexes.push(ColumnFamilyLogIndexPair {
applied_index: LogIndexSeqnoPair::new(0, 0),
flushed_index: LogIndexSeqnoPair::new(0, 0),
});
}
Self {
cf_indexes: Arc::new(RwLock::new(cf_indexes)),
last_flush_index: Arc::new(RwLock::new(LogIndexSeqnoPair::new(0, 0))),
}
}
/// 初始化:从 SST 文件读取最大的 log index(对应 C++ 的 Init)
///
/// **实现方案**:使用 TablePropertiesCollector(参考 C++ 版本)
/// 1. 在写入时通过 TablePropertiesCollector 收集 log_index 和 seqno
/// 2. 启动时遍历所有 SST 文件,从 TableProperties 读取最大的 log_index
/// 3. 如果 RocksDB Rust 绑定不支持 TableProperties,可以使用元数据文件记录
pub async fn init(&self, db: &DB, cf_handles: &[rocksdb::ColumnFamilyHandle]) -> Result<(), rocksdb::Error> {
let mut cf_indexes = self.cf_indexes.write().await;
for (i, cf_handle) in cf_handles.iter().enumerate() {
// 从 TableProperties 读取最大的 log index(生产级实现,对应 C++ 的 TablePropertiesCollector)
let mut max_log_index = 0u64;
let mut max_seqno = 0u64;
// 获取所有 SST 文件的 TableProperties(对应 C++ 的 GetPropertiesOfAllTables)
let props_collection = db.get_properties_of_all_tables(cf_handle)?;
for (_file_name, props) in props_collection.iter() {
// 从 user_collected_properties 读取 log_index 和 seqno
// 格式:"{log_index}/{seqno}"(对应 C++ 的 LogIndexTablePropertiesCollector)
if let Some(user_props) = props.user_collected_properties() {
if let Some(log_index_str) = user_props.get("log_index") {
if let Ok((log_idx, seqno)) = parse_log_index_seqno(log_index_str) {
if log_idx > max_log_index {
max_log_index = log_idx;
max_seqno = seqno;
}
}
}
}
}
// 设置初始值(对应 C++ 的 SetLogIndexSeqnoPair)
if let Some(cf) = cf_indexes.get_mut(i) {
cf.applied_index.set(max_log_index, max_seqno);
cf.flushed_index.set(max_log_index, max_seqno);
}
}
Ok(())
}
/// 解析 log_index 和 seqno(辅助函数,对应 C++ 的 LogIndexTablePropertiesCollector::ReadStatsFromTableProps)
fn parse_log_index_seqno(s: &str) -> Result<(u64, u64), std::num::ParseIntError> {
let parts: Vec<&str> = s.split('/').collect();
if parts.len() != 2 {
return Err(std::num::ParseIntError::from(std::io::Error::new(
std::io::ErrorKind::InvalidData,
"Invalid format",
)));
}
Ok((parts[0].parse()?, parts[1].parse()?))
}
/// 检查某个 CF 的日志是否已应用(对应 C++ 的 IsApplied)
pub async fn is_applied(&self, cf_id: usize, log_index: u64) -> bool {
let cf_indexes = self.cf_indexes.read().await;
if let Some(cf) = cf_indexes.get(cf_id) {
log_index < cf.applied_index.get_log_index()
} else {
false
}
}
/// 更新某个 CF 的 applied_index(对应 C++ 的 Update)
pub async fn update_applied_index(
&self,
cf_id: usize,
log_index: u64,
seqno: u64,
) {
let mut cf_indexes = self.cf_indexes.write().await;
if let Some(cf) = cf_indexes.get_mut(cf_id) {
// 检查是否需要更新 flushed_index(对应 C++ 的逻辑)
let last_flush = self.last_flush_index.read().await;
if cf.flushed_index.get_log_index() <= last_flush.get_log_index() &&
cf.flushed_index.get_log_index() == cf.applied_index.get_log_index()
{
let flush_log_index = cf.flushed_index.get_log_index()
.max(last_flush.get_log_index());
let flush_seqno = cf.flushed_index.get_seqno()
.max(last_flush.get_seqno());
cf.flushed_index.set(flush_log_index, flush_seqno);
}
// 更新 applied_index
cf.applied_index.set(log_index, seqno);
}
}
/// 设置某个 CF 的 flushed_index(对应 C++ 的 SetFlushedLogIndex)
pub async fn set_flushed_index(
&self,
cf_id: usize,
log_index: u64,
seqno: u64,
) {
let mut cf_indexes = self.cf_indexes.write().await;
if let Some(cf) = cf_indexes.get_mut(cf_id) {
let current_log = cf.flushed_index.get_log_index();
let current_seqno = cf.flushed_index.get_seqno();
cf.flushed_index.set(
current_log.max(log_index),
current_seqno.max(seqno),
);
}
}
/// 设置全局 flushed_index(对应 C++ 的 SetFlushedLogIndexGlobal)
pub async fn set_flushed_index_global(
&self,
log_index: u64,
seqno: u64,
) {
// 更新 last_flush_index
{
let mut last_flush = self.last_flush_index.write().await;
let current_log = last_flush.get_log_index();
let current_seqno = last_flush.get_seqno();
last_flush.set(
current_log.max(log_index),
current_seqno.max(seqno),
);
}
// 更新所有 CF 的 flushed_index(如果它们小于 last_flush_index)
let mut cf_indexes = self.cf_indexes.write().await;
let last_flush = self.last_flush_index.read().await;
for cf in cf_indexes.iter_mut() {
if cf.flushed_index.get_log_index() <= last_flush.get_log_index() {
let flush_log_index = cf.flushed_index.get_log_index()
.max(last_flush.get_log_index());
let flush_seqno = cf.flushed_index.get_seqno()
.max(last_flush.get_seqno());
cf.flushed_index.set(flush_log_index, flush_seqno);
}
}
}
/// 获取最小的 log index(对应 C++ 的 GetSmallestLogIndex)
pub async fn get_smallest_log_index(
&self,
exclude_cf: Option<usize>,
) -> SmallestIndexRes {
let cf_indexes = self.cf_indexes.read().await;
let mut smallest_applied_cf = 0;
let mut smallest_applied_index = LogIndex::MAX;
let mut smallest_flushed_cf = 0;
let mut smallest_flushed_index = LogIndex::MAX;
let mut smallest_flushed_seqno = SequenceNumber::MAX;
for (i, cf) in cf_indexes.iter().enumerate() {
if Some(i) == exclude_cf {
continue;
}
// 如果 flushed_index >= applied_index,跳过(对应 C++ 的逻辑)
if i != exclude_cf.unwrap_or(usize::MAX) &&
cf.flushed_index.get_log_index() >= cf.applied_index.get_log_index() {
continue;
}
let applied = cf.applied_index.get_log_index();
let flushed = cf.flushed_index.get_log_index();
let flushed_seqno = cf.flushed_index.get_seqno();
if applied < smallest_applied_index {
smallest_applied_index = applied;
smallest_applied_cf = i;
}
if flushed < smallest_flushed_index {
smallest_flushed_index = flushed;
smallest_flushed_seqno = flushed_seqno;
smallest_flushed_cf = i;
}
}
SmallestIndexRes {
smallest_applied_log_index_cf: smallest_applied_cf,
smallest_applied_log_index: smallest_applied_index,
smallest_flushed_log_index_cf: smallest_flushed_cf,
smallest_flushed_log_index: smallest_flushed_index,
smallest_flushed_seqno,
}
}
/// 获取全局最小 flushed_index(对应 C++ 的 GetLastFlushIndex)
pub async fn get_last_flush_index(&self) -> LogIndexSeqnoPair {
self.last_flush_index.read().await.clone()
}
/// 获取待 flush 的 gap(对应 C++ 的 GetPendingFlushGap)
pub async fn get_pending_flush_gap(&self) -> usize {
let cf_indexes = self.cf_indexes.read().await;
let mut indices = std::collections::BTreeSet::new();
for cf in cf_indexes.iter() {
indices.insert(cf.applied_index.get_log_index());
indices.insert(cf.flushed_index.get_log_index());
}
if indices.len() <= 1 {
return 0;
}
let first = indices.iter().next()
.ok_or_else(|| StorageError::IO {
source: StorageIOError::read_logs(&std::io::Error::new(
std::io::ErrorKind::InvalidData,
"Empty indices",
)),
})?;
let last = indices.iter().next_back()
.ok_or_else(|| StorageError::IO {
source: StorageIOError::read_logs(&std::io::Error::new(
std::io::ErrorKind::InvalidData,
"Empty indices",
)),
})?;
let first = *first;
let last = *last;
(last - first) as usize
}
}
/// LogIndex 和 SequenceNumber 收集器(对应 C++ 的 LogIndexAndSequenceCollector)
pub struct LogIndexAndSequenceCollector {
/// 采样掩码(对应 C++ 的 step_length_mask_)
step_length_mask: u64,
/// 日志索引和序列号对列表(对应 C++ 的 list_)
list: Arc<RwLock<VecDeque<(LogIndex, SequenceNumber)>>>,
/// 最大 gap(对应 C++ 的 max_gap_)
max_gap: Arc<AtomicU64>,
}
impl LogIndexAndSequenceCollector {
pub fn new(step_length_bit: u8) -> Self {
let step_length_mask = if step_length_bit > 0 {
(1 << step_length_bit) - 1
} else {
0
};
Self {
step_length_mask,
list: Arc::new(RwLock::new(VecDeque::new())),
max_gap: Arc::new(AtomicU64::new(1000)), // 默认 1000
}
}
/// 根据 seqno 查找对应的 log index(对应 C++ 的 FindAppliedLogIndex)
pub async fn find_applied_log_index(&self, seqno: SequenceNumber) -> LogIndex {
if seqno == 0 {
return 0;
}
let list = self.list.read().await;
if list.is_empty() || seqno < list[0].1 {
return 0;
}
if seqno >= list[list.len() - 1].1 {
return list[list.len() - 1].0;
}
// 二分查找
let mut left = 0;
let mut right = list.len();
while left < right {
let mid = (left + right) / 2;
if list[mid].1 <= seqno {
left = mid + 1;
} else {
right = mid;
}
}
if left > 0 {
list[left - 1].0
} else {
0
}
}
/// 更新收集器(对应 C++ 的 Update)
pub async fn update(&self, smallest_applied_log_index: LogIndex, smallest_flush_seqno: SequenceNumber) {
// 如果 step_length_mask > 0,采样以节省内存
if (smallest_applied_log_index & self.step_length_mask) == 0 {
let mut list = self.list.write().await;
list.push_back((smallest_applied_log_index, smallest_flush_seqno));
}
}
/// 清理过期的日志索引(对应 C++ 的 Purge)
pub async fn purge(&self, smallest_applied_log_index: LogIndex) {
let mut list = self.list.write().await;
if list.len() < 2 {
return;
}
// 删除所有小于等于 smallest_applied_log_index 的条目(保留至少 2 个)
while list.len() >= 2 && list.len() > 1 && list[1].0 <= smallest_applied_log_index {
list.pop_front();
}
}
/// 检查是否需要手动 flush(对应 C++ 的 IsFlushPending)
pub async fn is_flush_pending(&self) -> bool {
let list = self.list.read().await;
list.len() >= self.max_gap.load(Ordering::Acquire) as usize
}
/// 获取列表大小
pub async fn get_size(&self) -> usize {
let list = self.list.read().await;
list.len()
}
}文件:src/raft/src/log_store.rs
// src/raft/src/log_store.rs
use std::collections::BTreeMap;
use std::fmt::Debug;
use std::fs::{File, OpenOptions};
use std::io::{BufReader, BufWriter, Read, Seek, SeekFrom, Write};
use std::ops::RangeBounds;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex};
use byteorder::{BigEndian, ReadBytesExt, WriteBytesExt};
use openraft::storage::{LogFlushed, LogState};
use openraft::storage::RaftLogStorage;
use openraft::storage::RaftLogReader;
use openraft::{Entry, LogId, OptionalSend, StorageError, StorageIOError, Vote};
use openraft::AnyError;
use serde_json;
use crate::types::{TypeConfig, NodeId};
type StorageResult<T> = Result<T, StorageError<NodeId>>;
/// 段元数据(记录段的起始和结束 log_index)
#[derive(Debug, Clone, Serialize, Deserialize)]
struct SegmentMeta {
start_index: u64, // 段的起始 log_index
end_index: u64, // 段的结束 log_index(包含)
file_name: String, // 文件名
}
/// 段元数据集合
#[derive(Debug, Clone, Serialize, Deserialize)]
struct SegmentsMeta {
segments: Vec<SegmentMeta>,
}
/// 日志段起始索引转换为文件名(大端序,保证排序)
fn segment_start_to_filename(start_index: u64) -> String {
format!("{:016x}.log", start_index)
}
/// 文件名转换为段的起始索引
fn filename_to_segment_start(filename: &str) -> Option<u64> {
filename
.strip_suffix(".log")
.and_then(|s| u64::from_str_radix(s, 16).ok())
}
/// Raft 日志存储(使用文件系统,段模式,与 C++ 版本的 braft SegmentLogStorage 一致)
/// 日志存储在:db_path/db_id/_praft/log/
/// 元数据存储在:db_path/db_id/_praft/raft_meta/
///
/// **设计说明**:
/// - 使用段(Segment)模式:一个文件包含多个连续的 log_index 的 Entry
/// - 文件命名:基于段的起始 log_index:`{start_index:016x}.log`
/// - 文件大小限制:每个段文件最大 64MB(可配置),达到限制时创建新段
/// - 段元数据:`segments.json` 记录每个段的起始和结束 log_index
#[derive(Debug)]
pub struct KiwiLogStore {
/// 日志目录路径
log_dir: PathBuf,
/// 元数据目录路径
meta_dir: PathBuf,
/// 当前日志文件(用于追加)
current_log_file: Arc<Mutex<Option<BufWriter<File>>>>,
/// 当前段的起始 log_index
current_segment_start: Arc<Mutex<u64>>,
/// 当前段的结束 log_index(最后一个写入的 log_index)
current_segment_end: Arc<Mutex<u64>>,
/// 段元数据文件路径
segments_meta_path: PathBuf,
/// 段元数据(缓存)
segments_meta: Arc<Mutex<SegmentsMeta>>,
/// 最大文件大小(字节),默认 64MB
max_file_size: u64,
}
impl KiwiLogStore {
pub fn new<P: AsRef<Path>>(db_path: P, db_id: u32) -> StorageResult<Self> {
let base_path = db_path.as_ref().join(db_id.to_string()).join("_praft");
let log_dir = base_path.join("log");
let meta_dir = base_path.join("raft_meta");
let segments_meta_path = log_dir.join("segments.json");
// 创建目录
std::fs::create_dir_all(&log_dir).map_err(|e| StorageError::IO {
source: StorageIOError::read(&e),
})?;
std::fs::create_dir_all(&meta_dir).map_err(|e| StorageError::IO {
source: StorageIOError::read(&e),
})?;
// 加载段元数据
let segments_meta = Self::load_segments_meta(&segments_meta_path)?;
// 确定当前段的起始和结束索引
let (current_segment_start, current_segment_end) = if let Some(last_segment) = segments_meta.segments.last() {
(last_segment.start_index, last_segment.end_index)
} else {
(0, 0)
};
Ok(Self {
log_dir,
meta_dir,
current_log_file: Arc::new(Mutex::new(None)),
current_segment_start: Arc::new(Mutex::new(current_segment_start)),
current_segment_end: Arc::new(Mutex::new(current_segment_end)),
segments_meta_path,
segments_meta: Arc::new(Mutex::new(segments_meta)),
max_file_size: 64 * 1024 * 1024, // 默认 64MB
})
}
/// 加载段元数据(生产级实现,错误处理)
fn load_segments_meta(meta_path: &Path) -> StorageResult<SegmentsMeta> {
if !meta_path.exists() {
return Ok(SegmentsMeta { segments: Vec::new() });
}
let content = std::fs::read_to_string(meta_path).map_err(|e| StorageError::IO {
source: StorageIOError::read(&e),
})?;
serde_json::from_str(&content).map_err(|e| StorageError::IO {
source: StorageIOError::read(&std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("Failed to parse segments meta: {}", e),
)),
})
}
/// 保存段元数据(生产级实现,错误处理)
fn save_segments_meta(&self, meta: &SegmentsMeta) -> StorageResult<()> {
let content = serde_json::to_string(meta).map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("Failed to serialize segments meta: {}", e),
)),
})?;
std::fs::write(&self.segments_meta_path, content).map_err(|e| StorageError::IO {
source: StorageIOError::write(&e),
})?;
Ok(())
}
/// 根据 log_index 查找对应的段文件路径(生产级实现)
fn find_segment_file(&self, log_index: u64) -> StorageResult<Option<PathBuf>> {
let segments_meta = self.segments_meta.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::read(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock segments_meta: {}", e),
)),
})?;
// 查找包含该 log_index 的段
for segment in segments_meta.segments.iter() {
if log_index >= segment.start_index && log_index <= segment.end_index {
return Ok(Some(self.log_dir.join(&segment.file_name)));
}
}
Ok(None)
}
/// 查找最大的日志索引(从段元数据中查找)
fn find_max_log_index(&self) -> StorageResult<u64> {
let segments_meta = self.segments_meta.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::read(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock segments_meta: {}", e),
)),
})?;
Ok(segments_meta.segments
.iter()
.map(|s| s.end_index)
.max()
.unwrap_or(0))
}
/// 切换到新段(生产级实现,错误处理)
/// 当当前段文件达到大小限制时,创建新段
fn rotate_segment(&self, new_start_index: u64) -> StorageResult<()> {
// 刷新并关闭旧文件
{
let mut current_file = self.current_log_file.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_log_file: {}", e),
)),
})?;
if let Some(mut old_writer) = current_file.take() {
old_writer.flush().map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&e),
})?;
}
}
// 更新当前段的结束索引(如果旧段存在)
{
let mut segments_meta = self.segments_meta.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock segments_meta: {}", e),
)),
})?;
let old_end = *self.current_segment_end.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_segment_end: {}", e),
)),
})?;
// 更新最后一个段的结束索引
if let Some(last_segment) = segments_meta.segments.last_mut() {
last_segment.end_index = old_end;
}
}
// 创建新段
let new_file_name = segment_start_to_filename(new_start_index);
let new_file_path = self.log_dir.join(&new_file_name);
let file = OpenOptions::new()
.create(true)
.append(true)
.open(&new_file_path)
.map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&e),
})?;
// 更新当前段信息
{
let mut current_file = self.current_log_file.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_log_file: {}", e),
)),
})?;
*current_file = Some(BufWriter::new(file));
}
{
let mut segment_start = self.current_segment_start.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_segment_start: {}", e),
)),
})?;
*segment_start = new_start_index;
}
{
let mut segment_end = self.current_segment_end.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_segment_end: {}", e),
)),
})?;
*segment_end = new_start_index;
}
// 添加新段到元数据
{
let mut segments_meta = self.segments_meta.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock segments_meta: {}", e),
)),
})?;
segments_meta.segments.push(SegmentMeta {
start_index: new_start_index,
end_index: new_start_index,
file_name: new_file_name,
});
// 保存段元数据
self.save_segments_meta(&segments_meta)?;
}
Ok(())
}
/// 追加日志到当前文件(复用缓存句柄,生产级实现)
/// 如果文件达到大小限制,自动切换到新段
fn append_to_current_file(&self, entry: &Entry<TypeConfig>) -> StorageResult<()> {
let log_index = entry.log_id.index;
// 检查是否需要切换到新段
let need_new_segment = {
let mut current_file = self.current_log_file.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_log_file: {}", e),
)),
})?;
let segment_start = *self.current_segment_start.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_segment_start: {}", e),
)),
})?;
// 如果 log_index 不在当前段范围内,需要新段
if log_index < segment_start {
return Err(StorageError::IO {
source: StorageIOError::write_logs(&std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Log index {} is before current segment start {}", log_index, segment_start),
)),
});
}
// 检查文件大小
let file_too_large = current_file.as_ref()
.and_then(|w| w.get_ref().metadata().ok())
.map(|m| m.len() >= self.max_file_size)
.unwrap_or(false);
file_too_large
};
if need_new_segment {
// 切换到新段(新段的起始索引是当前 log_index)
self.rotate_segment(log_index)?;
}
// 确保文件已打开
{
let mut current_file = self.current_log_file.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_log_file: {}", e),
)),
})?;
if current_file.is_none() {
let segment_start = *self.current_segment_start.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_segment_start: {}", e),
)),
})?;
let file_path = self.log_dir.join(segment_start_to_filename(segment_start));
let file = OpenOptions::new()
.create(true)
.append(true)
.open(&file_path)
.map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&e),
})?;
*current_file = Some(BufWriter::new(file));
}
}
// 序列化 Entry
let entry_json = serde_json::to_vec(entry)
.map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("Failed to serialize entry: {}", e),
)),
})?;
// 写入数据(每行一个 Entry)
{
let mut current_file = self.current_log_file.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_log_file: {}", e),
)),
})?;
if let Some(writer) = current_file.as_mut() {
writer.write_all(&entry_json).map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&e),
})?;
writer.write_all(b"\n").map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&e),
})?;
}
}
// 更新当前段的结束索引
{
let mut segment_end = self.current_segment_end.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_segment_end: {}", e),
)),
})?;
*segment_end = log_index;
}
// 更新段元数据中的结束索引
{
let mut segments_meta = self.segments_meta.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock segments_meta: {}", e),
)),
})?;
if let Some(last_segment) = segments_meta.segments.last_mut() {
last_segment.end_index = log_index;
}
}
Ok(())
}
/// 读取段文件中的所有 Entry(生产级实现,错误处理)
fn read_segment_file(&self, file_path: &Path) -> StorageResult<Vec<Entry<TypeConfig>>> {
if !file_path.exists() {
return Ok(Vec::new());
}
let file = File::open(file_path).map_err(|e| StorageError::IO {
source: StorageIOError::read_logs(&e),
})?;
let mut reader = BufReader::new(file);
let mut entries = Vec::new();
let mut buffer = Vec::new();
// 读取文件内容
reader.read_to_end(&mut buffer).map_err(|e| StorageError::IO {
source: StorageIOError::read_logs(&e),
})?;
// 解析日志文件(生产级实现)
// 日志格式:每行一个 JSON 格式的 Entry(与 C++ 版本一致)
// 使用逐行解析,支持增量追加
let mut line_start = 0;
for (i, &byte) in buffer.iter().enumerate() {
if byte == b'\n' {
if i > line_start {
let line = &buffer[line_start..i];
if !line.is_empty() {
match serde_json::from_slice::<Entry<TypeConfig>>(line) {
Ok(entry) => entries.push(entry),
Err(e) => {
tracing::warn!("Failed to parse log entry at line {}: {}", entries.len(), e);
// 继续解析下一行,不中断
}
}
}
}
line_start = i + 1;
}
}
// 处理最后一行(如果没有换行符)
if line_start < buffer.len() {
let line = &buffer[line_start..];
if !line.is_empty() {
match serde_json::from_slice::<Entry<TypeConfig>>(line) {
Ok(entry) => entries.push(entry),
Err(e) => {
tracing::warn!("Failed to parse last log entry: {}", e);
}
}
}
}
Ok(entries)
}
/// 读取指定 log_index 的 Entry(生产级实现)
/// 先找到对应的段文件,然后读取并过滤
fn read_log_entry(&self, log_index: u64) -> StorageResult<Option<Entry<TypeConfig>>> {
let segment_file = self.find_segment_file(log_index)?;
if let Some(file_path) = segment_file {
let entries = self.read_segment_file(&file_path)?;
// 查找指定 log_index 的 Entry
Ok(entries.into_iter().find(|e| e.log_id.index == log_index))
} else {
Ok(None)
}
}
/// 写入元数据文件
fn write_meta_file(&self, key: &str, value: &[u8]) -> StorageResult<()> {
let file_path = self.meta_dir.join(key);
std::fs::write(&file_path, value).map_err(|e| StorageError::IO {
source: StorageIOError::write(&e),
})?;
Ok(())
}
/// 读取元数据文件
fn read_meta_file(&self, key: &str) -> StorageResult<Option<Vec<u8>>> {
let file_path = self.meta_dir.join(key);
if !file_path.exists() {
return Ok(None);
}
std::fs::read(&file_path).map_err(|e| StorageError::IO {
source: StorageIOError::read(&e),
}).map(Some)
}
fn get_last_purged_(&self) -> StorageResult<Option<LogId<NodeId>>> {
Ok(self
.read_meta_file("last_purged_log_id")?
.and_then(|v| serde_json::from_slice(&v).ok()))
}
fn set_last_purged_(&self, log_id: LogId<NodeId>) -> StorageResult<()> {
let data = serde_json::to_vec(&log_id)
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("Failed to serialize log_id: {}", e),
)),
})?;
self.write_meta_file("last_purged_log_id", &data)?;
Ok(())
}
fn set_committed_(&self, committed: &Option<LogId<NodeId>>) -> Result<(), StorageIOError<NodeId>> {
let json = serde_json::to_vec(committed)
.map_err(|e| StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("Failed to serialize committed: {}", e),
)))?;
self.write_meta_file("committed", &json)
.map_err(|e| StorageIOError::write(&e))?;
Ok(())
}
fn get_committed_(&self) -> StorageResult<Option<LogId<NodeId>>> {
Ok(self
.read_meta_file("committed")?
.and_then(|v| serde_json::from_slice(&v).ok()))
}
fn set_vote_(&self, vote: &Vote<NodeId>) -> StorageResult<()> {
let data = serde_json::to_vec(vote)
.map_err(|e| StorageError::IO {
source: StorageIOError::write_vote(&std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("Failed to serialize vote: {}", e),
)),
})?;
self.write_meta_file("vote", &data).map_err(|e| StorageError::IO {
source: StorageIOError::write_vote(&e),
})?;
Ok(())
}
fn get_vote_(&self) -> StorageResult<Option<Vote<NodeId>>> {
Ok(self
.read_meta_file("vote")?
.and_then(|v| serde_json::from_slice(&v).ok()))
}
}
/// RaftLogReader 实现
#[async_trait::async_trait]
impl RaftLogReader<TypeConfig> for KiwiLogStore {
async fn try_get_log_entries<RB: RangeBounds<u64> + Clone + Debug + OptionalSend>(
&mut self,
range: RB,
) -> StorageResult<Vec<Entry<TypeConfig>>> {
let mut entries = Vec::new();
// 获取范围边界
let start = match range.start_bound() {
std::ops::Bound::Included(x) => *x,
std::ops::Bound::Excluded(x) => *x + 1,
std::ops::Bound::Unbounded => 0,
};
let end = match range.end_bound() {
std::ops::Bound::Included(x) => Some(*x),
std::ops::Bound::Excluded(x) => Some(*x - 1),
std::ops::Bound::Unbounded => None,
};
// 从段元数据中找到所有相关的段文件
let segments_meta = self.segments_meta.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::read(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock segments_meta: {}", e),
)),
})?;
for segment in segments_meta.segments.iter() {
// 检查段是否与范围有交集
let segment_start = segment.start_index;
let segment_end = segment.end_index;
let range_start = start;
let range_end = end.unwrap_or(u64::MAX);
if segment_end >= range_start && segment_start <= range_end {
// 读取段文件
let file_path = self.log_dir.join(&segment.file_name);
let segment_entries = self.read_segment_file(&file_path)?;
// 过滤出范围内的 Entry
for entry in segment_entries {
if range.contains(&entry.log_id.index) {
entries.push(entry);
}
}
}
}
// 按 log_index 排序(确保顺序正确)
entries.sort_by_key(|e| e.log_id.index);
Ok(entries)
}
}
/// RaftLogStorage 实现
#[async_trait::async_trait]
impl RaftLogStorage<TypeConfig> for KiwiLogStore {
type LogReader = Self;
async fn get_log_state(&mut self) -> StorageResult<LogState<TypeConfig>> {
// 获取最后一个日志条目(从段元数据中查找)
let max_index = self.find_max_log_index()?;
let last_log_id = if max_index > 0 {
self.read_log_entry(max_index)?
} else {
None
};
let last_purged_log_id = self.get_last_purged_()?;
let last_log_id = match last_log_id {
None => last_purged_log_id.clone(),
Some(entry) => Some(entry.log_id),
};
Ok(LogState {
last_purged_log_id,
last_log_id,
})
}
async fn save_committed(&mut self, committed: Option<LogId<NodeId>>) -> Result<(), StorageError<NodeId>> {
self.set_committed_(&committed)?;
Ok(())
}
async fn read_committed(&mut self) -> Result<Option<LogId<NodeId>>, StorageError<NodeId>> {
self.get_committed_()
}
async fn save_vote(&mut self, vote: &Vote<NodeId>) -> Result<(), StorageError<NodeId>> {
self.set_vote_(vote)
}
async fn read_vote(&mut self) -> Result<Option<Vote<NodeId>>, StorageError<NodeId>> {
self.get_vote_()
}
async fn append<I>(
&mut self,
entries: I,
callback: LogFlushed<TypeConfig>,
) -> StorageResult<()>
where
I: IntoIterator<Item = Entry<TypeConfig>> + Send,
I::IntoIter: Send,
{
// 追加日志条目到文件(生产级实现,段模式,类似 C++ 的 braft SegmentLogStorage)
// 使用 append_to_current_file 复用文件句柄,自动管理段切换
for entry in entries {
self.append_to_current_file(&entry)?;
}
// 刷新文件(确保数据写入磁盘)
{
let mut current_file = self.current_log_file.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_log_file: {}", e),
)),
})?;
if let Some(writer) = current_file.as_mut() {
writer.flush().map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&e),
})?;
}
}
// 保存段元数据(确保元数据持久化)
{
let segments_meta = self.segments_meta.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock segments_meta: {}", e),
)),
})?;
self.save_segments_meta(&segments_meta)?;
}
// 通知日志写入完成
callback.log_io_completed(Ok(()));
Ok(())
}
async fn truncate(&mut self, log_id: LogId<NodeId>) -> StorageResult<()> {
// 截断日志(删除 log_id 之后的所有段文件,生产级实现)
let mut segments_meta = self.segments_meta.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock segments_meta: {}", e),
)),
})?;
// 找到需要删除的段
let mut segments_to_remove = Vec::new();
for (i, segment) in segments_meta.segments.iter().enumerate() {
if segment.start_index > log_id.index {
segments_to_remove.push(i);
} else if segment.start_index <= log_id.index && segment.end_index > log_id.index {
// 部分截断:更新段的结束索引
segment.end_index = log_id.index;
}
}
// 删除段文件(从后往前删除,避免索引变化)
for &i in segments_to_remove.iter().rev() {
let segment = &segments_meta.segments[i];
let file_path = self.log_dir.join(&segment.file_name);
if file_path.exists() {
std::fs::remove_file(&file_path).map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&e),
})?;
}
}
// 从元数据中移除段
for &i in segments_to_remove.iter().rev() {
segments_meta.segments.remove(i);
}
// 更新当前段信息
if let Some(last_segment) = segments_meta.segments.last() {
*self.current_segment_start.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_segment_start: {}", e),
)),
})? = last_segment.start_index;
*self.current_segment_end.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock current_segment_end: {}", e),
)),
})? = last_segment.end_index;
}
// 保存段元数据
self.save_segments_meta(&segments_meta)?;
Ok(())
}
async fn purge(&mut self, log_id: LogId<NodeId>) -> Result<(), StorageError<NodeId>> {
// 清理日志(删除 log_id 之前的所有段文件,生产级实现)
self.set_last_purged_(log_id)?;
let mut segments_meta = self.segments_meta.lock()
.map_err(|e| StorageError::IO {
source: StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to lock segments_meta: {}", e),
)),
})?;
// 找到需要删除的段(所有结束索引 <= log_id.index 的段)
let mut segments_to_remove = Vec::new();
for (i, segment) in segments_meta.segments.iter().enumerate() {
if segment.end_index <= log_id.index {
segments_to_remove.push(i);
} else if segment.start_index <= log_id.index {
// 部分清理:更新段的起始索引
segment.start_index = log_id.index + 1;
}
}
// 删除段文件(从前往后删除)
for &i in segments_to_remove.iter() {
let segment = &segments_meta.segments[i];
let file_path = self.log_dir.join(&segment.file_name);
if file_path.exists() {
std::fs::remove_file(&file_path).map_err(|e| StorageError::IO {
source: StorageIOError::write_logs(&e),
})?;
}
}
// 从元数据中移除段(从后往前删除,避免索引变化)
for &i in segments_to_remove.iter().rev() {
segments_meta.segments.remove(i);
}
// 保存段元数据
self.save_segments_meta(&segments_meta)?;
Ok(())
}
async fn get_log_reader(&mut self) -> Self::LogReader {
// 创建新的实例用于读取(生产级实现,错误处理)
let parent = self.log_dir.parent()
.and_then(|p| p.parent())
.ok_or_else(|| StorageError::IO {
source: StorageIOError::read_logs(&std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Invalid log_dir path",
)),
})?;
// 从路径中提取 db_id
let db_id_str = parent.file_name()
.and_then(|n| n.to_str())
.ok_or_else(|| StorageError::IO {
source: StorageIOError::read_logs(&std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Invalid db_id in path",
)),
})?;
let db_id = db_id_str.parse::<u32>()
.map_err(|e| StorageError::IO {
source: StorageIOError::read_logs(&std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Failed to parse db_id: {}", e),
)),
})?;
KiwiLogStore::new(parent.parent().ok_or_else(|| StorageError::IO {
source: StorageIOError::read_logs(&std::io::Error::new(
std::io::ErrorKind::InvalidInput,
"Invalid parent path",
)),
})?, db_id)
.map_err(|e| {
tracing::error!("Failed to create log reader: {:?}", e);
e
})?
}
}
impl Clone for KiwiLogStore {
fn clone(&self) -> Self {
// 克隆实现(生产级实现)
// 注意:文件句柄不共享,每次克隆都创建新的句柄
// 这与 C++ 版本一致,避免文件句柄共享带来的并发问题
let segment_start = self.current_segment_start.lock()
.map(|guard| *guard)
.unwrap_or(0); // 如果锁失败,使用默认值 0
let segment_end = self.current_segment_end.lock()
.map(|guard| *guard)
.unwrap_or(0); // 如果锁失败,使用默认值 0
// 重新加载段元数据(确保一致性)
let segments_meta = Self::load_segments_meta(&self.segments_meta_path)
.unwrap_or_else(|_| SegmentsMeta { segments: Vec::new() });
Self {
log_dir: self.log_dir.clone(),
meta_dir: self.meta_dir.clone(),
current_log_file: Arc::new(Mutex::new(None)), // 不共享文件句柄
current_segment_start: Arc::new(Mutex::new(segment_start)),
current_segment_end: Arc::new(Mutex::new(segment_end)),
segments_meta_path: self.segments_meta_path.clone(),
segments_meta: Arc::new(Mutex::new(segments_meta)),
max_file_size: self.max_file_size,
}
}
}注意:
- 段模式:使用段(Segment)模式组织日志文件,一个文件包含多个连续的 log_index 的 Entry
-
文件命名:基于段的起始 log_index:
{start_index:016x}.log,便于排序和查找 - 文件大小限制:每个段文件最大 64MB(可配置),达到限制时自动创建新段
-
段元数据:
segments.json记录每个段的起始和结束 log_index,用于快速查找 - 日志格式:每行一个 JSON 格式的 Entry,便于追加和读取
- 元数据存储:使用独立的文件存储 vote、committed、last_purged_log_id
-
目录结构:与 C++ 版本一致,
db_path/db_id/_praft/log/和db_path/db_id/_praft/raft_meta/ -
错误处理:所有操作都使用
Result类型,不使用unwrap(),确保生产级可靠性
文件:src/raft/src/state_machine.rs
// src/raft/src/state_machine.rs
use std::io::Cursor;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use tokio::sync::RwLock;
use openraft::storage::{RaftStateMachine, RaftSnapshotBuilder};
use openraft::{Entry, EntryPayload, LogId, OptionalSend, Snapshot, SnapshotMeta, StorageError, StoredMembership};
use rocksdb::DB;
use serde::{Deserialize, Serialize};
use crate::types::{TypeConfig, NodeId, Node, Request, Response, SnapshotData, Binlog, BinlogEntry, OperateType};
use crate::log_index::{LogIndexOfColumnFamilies, LogIndexAndSequenceCollector};
use storage::db::DB; // 使用 DB 包装结构体
use storage::StorageOptions;
use std::path::PathBuf;
use std::sync::Arc;
use tokio::fs::File;
use prost::Message; // 用于 protobuf 序列化/反序列化
type StorageResult<T> = Result<T, StorageError<NodeId>>;
/// 快照数据(持久化)
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct StoredSnapshot {
pub meta: SnapshotMeta<NodeId, Node>,
pub data: Vec<u8>,
}
/// Kiwi 状态机(操作级日志,类似 C++ 的 StateMachine)
pub struct KiwiStateMachine {
/// DB 包装结构体(对应 C++ 的 DB,包含 storage_mutex)
db: Arc<DB>,
/// CF 索引管理
log_indexes: Arc<LogIndexOfColumnFamilies>,
/// LogIndex 收集器
log_collector: Arc<LogIndexAndSequenceCollector>,
/// 最后应用的日志 ID(全局)
last_applied_log_id: Arc<RwLock<Option<LogId<NodeId>>>>,
/// 最后成员配置
last_membership: Arc<RwLock<StoredMembership<NodeId, Node>>>,
/// Storage 选项(用于重启数据库)
storage_options: Arc<StorageOptions>,
/// 快照索引
snapshot_idx: Arc<RwLock<u64>>,
/// 快照目录
snapshot_dir: PathBuf,
}
impl KiwiStateMachine {
pub fn new(
db: Arc<DB>,
log_indexes: Arc<LogIndexOfColumnFamilies>,
log_collector: Arc<LogIndexAndSequenceCollector>,
storage_options: Arc<StorageOptions>,
snapshot_dir: PathBuf,
) -> Self {
Self {
db,
log_indexes,
log_collector,
last_applied_log_id: Arc::new(RwLock::new(None)),
last_membership: Arc::new(RwLock::new(StoredMembership::default())),
storage_options,
snapshot_idx: Arc::new(RwLock::new(0)),
snapshot_dir,
}
}
/// 获取 Storage(对应 C++ 的 GetStorage)
pub async fn get_storage(&self) -> Option<Arc<Storage>> {
self.db.get_storage().await
}
/// 应用 Binlog(完整实现,对应 C++ 的 OnBinlogWrite)
async fn apply_binlog(
&self,
binlog: &Binlog,
log_idx: u64,
) -> StorageResult<Response> {
// 获取对应的 Redis 实例(对应 C++ 的 insts_[log.slot_idx()])
let slot_idx = binlog.slot_idx as usize;
// 获取 Storage(使用共享锁)
let storage = self.get_storage().await
.ok_or_else(|| StorageError::IO {
source: openraft::StorageIOError::read_logs(&std::io::Error::new(
std::io::ErrorKind::NotFound,
"Storage not found",
)),
})?;
// 获取对应的 Redis 实例
let instance = storage.insts.get(slot_idx)
.ok_or_else(|| StorageError::IO {
source: openraft::StorageIOError::read_logs(&std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Invalid slot_idx: {}", slot_idx),
)),
})?;
// 创建 WriteBatch(对应 C++ 的 rocksdb::WriteBatch batch)
let mut batch_ops = Vec::new();
let mut is_finished_start = true;
// 获取当前 SequenceNumber(生产级实现,对应 C++ 的 GetLatestSequenceNumber)
let mut seqno = instance.get_latest_sequence_number()
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::read_logs(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to get latest sequence number: {}", e),
)),
})?;
// 遍历 Binlog entries(对应 C++ 的 for (const auto& entry : log.entries()))
for entry in &binlog.entries {
let cf_idx = entry.cf_idx as usize;
// 检查是否已应用(生产级实现,对应 C++ 的 IsRestarting && IsApplied)
if instance.is_restarting() && self.log_indexes.is_applied(cf_idx, log_idx).await {
// 如果已应用,跳过(重启恢复时避免重复应用)
tracing::warn!("Log {} has been applied for CF {}", log_idx, cf_idx);
is_finished_start = false;
continue;
}
// 根据操作类型执行(对应 C++ 的 switch (entry.op_type()))
// 注意:protobuf 生成的 OperateType 是 i32,需要转换为 Rust enum
let op_type = OperateType::from_i32(entry.op_type)
.ok_or_else(|| StorageError::IO {
source: StorageIOError::write(&format!("Invalid OperateType: {}", entry.op_type)),
})?;
match op_type {
OperateType::KPut => {
// PUT 操作(对应 C++ 的 case OperateType::kPut)
if let Some(value) = &entry.value {
batch_ops.push((cf_idx, BatchOp::Put {
key: entry.key.clone(),
value: value.clone(),
}));
}
}
OperateType::KDelete => {
// DELETE 操作(对应 C++ 的 case OperateType::kDelete)
batch_ops.push((cf_idx, BatchOp::Delete {
key: entry.key.clone(),
}));
}
OperateType::KNoOperate => {
// 无操作,跳过
continue;
}
}
// 更新 CF 的 applied_index(对应 C++ 的 UpdateAppliedLogIndexOfColumnFamily)
seqno += 1; // 对应 C++ 的 ++seqno
self.log_indexes.update_applied_index(cf_idx, log_idx, seqno).await;
}
// 如果重启阶段完成,标记为完成(生产级实现,对应 C++ 的 StartingPhaseEnd)
if instance.is_restarting() && is_finished_start {
tracing::info!("Redis {} finished start phase", slot_idx);
instance.starting_phase_end();
}
// 批量执行操作(生产级实现,对应 C++ 的 db->Write(batch))
// 使用 WriteBatch 批量写入,提高性能
let mut write_batch = rocksdb::WriteBatch::default();
let cf_handles = instance.get_column_family_handles();
for (cf_idx, op) in batch_ops {
let cf_handle = cf_handles.get(cf_idx)
.ok_or_else(|| StorageError::IO {
source: openraft::StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Invalid CF index: {}", cf_idx),
)),
})?;
match op {
BatchOp::Put { key, value } => {
write_batch.put_cf(cf_handle, &key, &value);
}
BatchOp::Delete { key } => {
write_batch.delete_cf(cf_handle, &key);
}
}
}
// 获取写入后的第一个 SequenceNumber(对应 C++ 的 GetLatestSequenceNumber() + 1)
let first_seqno = instance.get_latest_sequence_number()
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::write(&e),
})? + 1;
// 写入 RocksDB(对应 C++ 的 Write(GetWriteOptions(), &batch))
let write_opts = instance.get_write_options();
instance.get_db().write_opt(&write_batch, &write_opts)
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::write(&std::io::Error::new(
std::io::ErrorKind::Other,
format!("Failed to write batch: {}", e),
)),
})?;
// 更新 LogIndex 收集器(对应 C++ 的 UpdateLogIndex)
instance.update_log_index(log_idx, first_seqno).await;
Ok(Response {
success: true,
error: None,
})
}
}
/// 批量操作类型(内部使用)
#[derive(Debug, Clone)]
enum BatchOp {
Put { key: Vec<u8>, value: Vec<u8> },
Delete { key: Vec<u8> },
}
/// RaftSnapshotBuilder 实现
#[async_trait::async_trait]
impl RaftSnapshotBuilder<TypeConfig> for KiwiStateMachine {
async fn build_snapshot(&mut self) -> StorageResult<Snapshot<TypeConfig>> {
let last_applied_log = self.last_applied_log_id.read().await.clone();
let last_membership = self.last_membership.read().await.clone();
// 获取全局最小 flushed_index(对应 C++ 的 GetSmallestFlushedLogIndex)
let last_flush_index = self.log_indexes.get_last_flush_index().await;
let snapshot_log_index = last_flush_index.get_log_index();
// 生成快照 ID(对应 C++ 的 snapshot_id)
let snapshot_id = if let Some(last) = last_applied_log.as_ref() {
let mut idx = self.snapshot_idx.write().await;
*idx += 1;
format!("{}-{}-{}", last.leader_id, last.index, *idx)
} else {
let mut idx = self.snapshot_idx.write().await;
*idx += 1;
format!("--{}", *idx)
};
// 创建快照目录(对应 C++ 的 snapshot_path)
let snapshot_path = self.snapshot_dir.join(&snapshot_id);
std::fs::create_dir_all(&snapshot_path)
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::write_snapshot(
Some(snapshot_id.clone()),
&e,
),
})?;
// 创建 Checkpoint(对应 C++ 的 CreateCheckpoint)
// 使用共享锁,与 C++ 版本一致
self.db.create_checkpoint(&snapshot_path, true)
.await
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::write_snapshot(
Some(snapshot_id.clone()),
&e,
),
})?;
// 打开快照文件(用于分块传输)
// openraft 会通过 install_snapshot RPC 分块读取这个文件
let snapshot_file = File::open(&snapshot_path)
.await
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::write_snapshot(
Some(snapshot_id.clone()),
&e,
),
})?;
// 获取 leader_id
let leader_id = last_applied_log.as_ref()
.map(|log_id| log_id.leader_id)
.unwrap_or(0);
let meta = SnapshotMeta {
last_log_id: Some(LogId {
leader_id,
index: snapshot_log_index,
}),
last_membership,
snapshot_id: snapshot_id.clone(),
};
// 保存快照元数据(生产级实现,对应 C++ 的快照元数据保存)
let stored_snapshot = StoredSnapshot {
meta: meta.clone(),
data: snapshot_path.to_string_lossy().into_owned().into_bytes(),
};
// 保存到文件系统(对应 C++ 的快照元数据文件)
let snapshot_meta_path = self.snapshot_dir.join(format!("{}.meta", snapshot_id));
let meta_json = serde_json::to_string(&stored_snapshot)
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::write_snapshot(
Some(snapshot_id.clone()),
&std::io::Error::new(std::io::ErrorKind::Other, e),
),
})?;
tokio::fs::write(&snapshot_meta_path, meta_json).await
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::write_snapshot(
Some(snapshot_id.clone()),
&e,
),
})?;
Ok(Snapshot {
meta,
snapshot: Box::new(snapshot_file),
})
}
}
/// RaftStateMachine 实现(操作级日志应用,类似 C++ 的 on_apply)
#[async_trait::async_trait]
impl RaftStateMachine<TypeConfig> for KiwiStateMachine {
type SnapshotBuilder = Self;
async fn applied_state(
&mut self,
) -> StorageResult<(Option<LogId<NodeId>>, StoredMembership<NodeId, Node>)> {
Ok((
self.last_applied_log_id.read().await.clone(),
self.last_membership.read().await.clone(),
))
}
async fn apply<I>(&mut self, entries: I) -> StorageResult<Vec<Response>>
where
I: IntoIterator<Item = Entry<TypeConfig>> + OptionalSend,
I::IntoIter: OptionalSend,
{
let entries: Vec<_> = entries.into_iter().collect();
let mut responses = Vec::with_capacity(entries.len());
for ent in entries {
let log_index = ent.log_id.index;
// 更新 last_applied_log_id
{
let mut last = self.last_applied_log_id.write().await;
*last = Some(ent.log_id.clone());
}
match ent.payload {
EntryPayload::Blank => {
responses.push(Response {
success: true,
error: None,
});
}
EntryPayload::Normal(binlog) => {
// 应用 Binlog(操作级日志,类似 C++ 的 OnBinlogWrite)
// EntryPayload::Normal 直接包含 Binlog 类型(不是 bytes)
// openraft 在读取日志时已经通过 serde 反序列化了整个 Entry
let result = self.apply_binlog(&binlog, log_index).await?;
responses.push(result);
}
EntryPayload::Membership(mem) => {
// 成员变更
let mut membership = self.last_membership.write().await;
*membership = StoredMembership::new(Some(ent.log_id.clone()), mem);
responses.push(Response {
success: true,
error: None,
});
}
}
}
Ok(responses)
}
async fn get_snapshot_builder(&mut self) -> Self::SnapshotBuilder {
self.clone()
}
async fn begin_receiving_snapshot(&mut self) -> StorageResult<Box<File>> {
// 创建临时文件用于接收快照数据(分块传输)
// openraft 会通过 install_snapshot RPC 分块写入这个文件
// 对应 C++ 版本的临时文件创建
let temp_file = std::env::temp_dir()
.join(format!("snapshot_receiving_{}", std::process::id()));
let file = File::create(&temp_file)
.await
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::read_snapshot(
None,
&e,
),
})?;
Ok(Box::new(file))
}
async fn install_snapshot(
&mut self,
meta: &SnapshotMeta<NodeId, Node>,
snapshot: Box<SnapshotData>,
) -> StorageResult<()> {
// 从快照文件恢复状态机(对应 C++ 的 LoadDBFromCheckpoint)
// snapshot 是一个 File 句柄,openraft 已经通过分块传输写入了快照数据
// 获取快照文件路径(从临时文件)
// 注意:这里需要保存快照文件路径,因为 File 句柄关闭后无法再访问
// 实际实现中,应该在 begin_receiving_snapshot 中保存路径
// 创建快照目录
let snapshot_path = self.snapshot_dir.join(&meta.snapshot_id);
std::fs::create_dir_all(&snapshot_path)
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::read_snapshot(
Some(meta.snapshot_id.clone()),
&e,
),
})?;
// 关闭文件句柄,确保数据已写入
drop(snapshot);
// 从 Checkpoint 加载数据库(对应 C++ 的 LoadDBFromCheckpoint)
// 使用独占锁,与 C++ 版本一致
self.db.load_db_from_checkpoint(
&snapshot_path,
self.storage_options.clone(),
true, // sync = true
)
.await
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::read_snapshot(
Some(meta.snapshot_id.clone()),
&e,
),
})?;
// 更新 last_applied_log_id
if let Some(last_log_id) = meta.last_log_id {
let mut last = self.last_applied_log_id.write().await;
*last = Some(last_log_id);
}
// 更新 last_membership
{
let mut membership = self.last_membership.write().await;
*membership = StoredMembership::new(
meta.last_log_id,
meta.last_membership.clone(),
);
}
Ok(())
}
async fn get_current_snapshot(&mut self) -> StorageResult<Option<Snapshot<TypeConfig>>> {
// 从持久化存储加载快照元数据(生产级实现,对应 C++ 的快照加载)
// 查找最新的快照元数据文件
let mut snapshot_files = tokio::fs::read_dir(&self.snapshot_dir).await
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::read_snapshot(&e),
})?;
let mut meta_files = Vec::new();
while let Some(entry) = snapshot_files.next_entry().await
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::read_snapshot(&e),
})? {
let path = entry.path();
if path.extension().and_then(|s| s.to_str()) == Some("meta") {
meta_files.push(path);
}
}
// 按文件名降序排序,获取最新的快照
meta_files.sort_by(|a, b| b.cmp(a));
// 读取最新的快照元数据
if let Some(meta_path) = meta_files.first() {
let meta_json = tokio::fs::read_to_string(meta_path).await
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::read_snapshot(&e),
})?;
let stored_snapshot: StoredSnapshot = serde_json::from_str(&meta_json)
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::read_snapshot(
&std::io::Error::new(std::io::ErrorKind::InvalidData, e),
),
})?;
// 打开快照文件
let snapshot_path = std::path::PathBuf::from(String::from_utf8_lossy(&stored_snapshot.data));
let file = File::open(&snapshot_path).await
.map_err(|e| StorageError::IO {
source: openraft::StorageIOError::read_snapshot(&e),
})?;
Ok(Some(Snapshot {
meta: stored_snapshot.meta,
snapshot: Box::new(file),
}))
} else {
Ok(None)
}
}
}
impl Clone for KiwiStateMachine {
fn clone(&self) -> Self {
Self {
db: self.db.clone(),
log_indexes: self.log_indexes.clone(),
log_collector: self.log_collector.clone(),
last_applied_log_id: self.last_applied_log_id.clone(),
last_membership: self.last_membership.clone(),
storage_options: self.storage_options.clone(),
snapshot_idx: self.snapshot_idx.clone(),
snapshot_dir: self.snapshot_dir.clone(),
}
}
}文件:src/raft/src/flush_manager.rs
// src/raft/src/flush_manager.rs
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::Arc;
use rocksdb::{DB, EventListener, FlushJobInfo};
use tokio::sync::RwLock;
use crate::log_index::{LogIndexOfColumnFamilies, LogIndexAndSequenceCollector};
/// Flush 完成回调函数类型
pub type FlushCallback = Box<dyn Fn(u64, bool) + Send + Sync>;
/// Flush 事件监听器(对应 C++ 的 LogIndexAndSequenceCollectorPurger)
pub struct FlushEventListener {
/// ColumnFamily handles
cf_handles: Arc<Vec<rocksdb::ColumnFamilyHandle>>,
/// LogIndex 收集器
collector: Arc<LogIndexAndSequenceCollector>,
/// CF 索引管理
cf_indexes: Arc<LogIndexOfColumnFamilies>,
/// Flush 计数器
count: Arc<AtomicU64>,
/// 手动 flush 的 CF ID
manual_flushing_cf: Arc<AtomicUsize>,
/// Flush 完成回调
callback: Arc<FlushCallback>,
}
impl FlushEventListener {
pub fn new(
cf_handles: Arc<Vec<rocksdb::ColumnFamilyHandle>>,
collector: Arc<LogIndexAndSequenceCollector>,
cf_indexes: Arc<LogIndexOfColumnFamilies>,
callback: FlushCallback,
) -> Self {
Self {
cf_handles,
collector,
cf_indexes,
count: Arc::new(AtomicU64::new(0)),
manual_flushing_cf: Arc::new(AtomicUsize::new(usize::MAX)),
callback: Arc::new(callback),
}
}
}
impl EventListener for FlushEventListener {
fn on_flush_completed(&self, _db: &DB, flush_job_info: &FlushJobInfo) {
let cf_id = flush_job_info.cf_id;
let largest_seqno = flush_job_info.largest_seqno;
// 根据 largest_seqno 查找对应的 log index
let rt = tokio::runtime::Handle::current();
let log_index = rt.block_on(self.collector.find_applied_log_index(largest_seqno));
// 更新该 CF 的 flushed_index(对应 C++ 的 SetFlushedLogIndex)
rt.block_on(self.cf_indexes.set_flushed_index(cf_id, log_index, largest_seqno));
// 获取最小的 log index(对应 C++ 的 GetSmallestLogIndex)
let smallest = rt.block_on(self.cf_indexes.get_smallest_log_index(Some(cf_id)));
// 清理过期的日志索引(对应 C++ 的 Purge)
rt.block_on(self.collector.purge(smallest.smallest_applied_log_index));
// 更新全局 flushed_index(对应 C++ 的 SetFlushedLogIndexGlobal)
if smallest.smallest_flushed_log_index_cf != usize::MAX {
rt.block_on(self.cf_indexes.set_flushed_index_global(
smallest.smallest_flushed_log_index,
smallest.smallest_flushed_seqno,
));
}
// 每 10 次 flush 触发一次快照(对应 C++ 的逻辑)
let count = self.count.fetch_add(1, Ordering::Relaxed);
if count % 10 == 0 {
tracing::info!("do snapshot after flush: {}", smallest.smallest_flushed_log_index);
(self.callback)(smallest.smallest_flushed_log_index, false);
}
// 如果当前 flush 的是手动触发的 CF,清除标记
if cf_id == self.manual_flushing_cf.load(Ordering::Acquire) {
self.manual_flushing_cf.store(usize::MAX, Ordering::Release);
}
// 检查是否需要手动 flush 落后的 CF(对应 C++ 的逻辑)
let flushing_cf = self.manual_flushing_cf.load(Ordering::Acquire);
if flushing_cf != usize::MAX {
return; // 已经有 CF 在 flush
}
if !rt.block_on(self.collector.is_flush_pending()) {
return; // 不需要 flush
}
// 尝试设置手动 flush 的 CF
let expected = usize::MAX;
if self.manual_flushing_cf.compare_exchange(
expected,
smallest.smallest_flushed_log_index_cf,
Ordering::Acquire,
Ordering::Relaxed,
).is_ok() {
// 触发手动 flush(对应 C++ 的 db->Flush)
if let Some(cf_handle) = self.cf_handles.get(smallest.smallest_flushed_log_index_cf) {
let mut flush_options = rocksdb::FlushOptions::default();
flush_options.set_wait(false);
if let Err(e) = _db.flush_opt(&flush_options, cf_handle) {
tracing::error!("Failed to flush CF {}: {}", smallest.smallest_flushed_log_index_cf, e);
self.manual_flushing_cf.store(usize::MAX, Ordering::Release);
}
}
}
}
}文件:src/raft/src/snapshot.rs
// src/raft/src/snapshot.rs
use std::path::{Path, PathBuf};
use std::sync::Arc;
use tokio::sync::RwLock;
use rocksdb::DB;
use openraft::{LogId, NodeId};
use crate::log_index::LogIndexOfColumnFamilies;
use storage::storage::Storage;
/// 快照管理器(对应 C++ 的 snapshot 逻辑)
pub struct SnapshotManager {
/// Storage 实例
storage: Arc<Storage>,
/// DB 实例
db: Arc<DB>,
/// CF 索引管理
log_indexes: Arc<LogIndexOfColumnFamilies>,
/// 快照路径
snapshot_path: PathBuf,
}
impl SnapshotManager {
pub fn new(
storage: Arc<Storage>,
db: Arc<DB>,
log_indexes: Arc<LogIndexOfColumnFamilies>,
snapshot_path: PathBuf,
) -> Self {
Self {
storage,
db,
log_indexes,
snapshot_path,
}
}
/// 创建快照(对应 C++ 的 snapshot 生成逻辑)
pub async fn create_snapshot(&self, db_id: u32) -> Result<LogId<NodeId>, Box<dyn std::error::Error>> {
// 获取全局最小 flushed_index(对应 C++ 的 GetSmallestFlushedLogIndex)
let last_flush_index = self.log_indexes.get_last_flush_index().await;
let snapshot_log_index = last_flush_index.get_log_index();
tracing::info!("Start to generate snapshot for DB {} at log index {}", db_id, snapshot_log_index);
// 创建 checkpoint(对应 C++ 的 CreateCheckpoint)
let checkpoint_path = self.snapshot_path.join(format!("db_{}", db_id));
std::fs::create_dir_all(&checkpoint_path)?;
// TODO: 创建 RocksDB checkpoint
// let checkpoint = rocksdb::checkpoint::Checkpoint::new(&self.db)?;
// checkpoint.create_checkpoint(&checkpoint_path)?;
// 更新快照元数据(对应 C++ 的 snapshot meta 更新)
// TODO: 实现快照元数据的保存
tracing::info!("Succeed to generate snapshot for DB {} at log index {}", db_id, snapshot_log_index);
Ok(LogId {
leader_id: 0, // TODO: 从 Raft 获取
index: snapshot_log_index,
})
}
/// 加载快照(对应 C++ 的快照加载逻辑)
pub async fn load_snapshot(&self, db_id: u32, snapshot_path: &Path) -> Result<(), Box<dyn std::error::Error>> {
tracing::info!("Load snapshot for DB {} from {}", db_id, snapshot_path.display());
// TODO: 从快照路径加载 RocksDB checkpoint
// 1. 关闭当前 DB
// 2. 复制快照文件到数据目录
// 3. 重新打开 DB
Ok(())
}
/// 获取快照的 last_log_index(对应 C++ 的 GetSmallestFlushedLogIndex)
pub async fn get_snapshot_last_log_index(&self) -> u64 {
let last_flush_index = self.log_indexes.get_last_flush_index().await;
last_flush_index.get_log_index()
}
}文件:src/raft/src/network.rs
// src/raft/src/network.rs
use std::collections::HashMap;
use std::sync::Arc;
use openraft::network::{RaftNetwork, RaftNetworkFactory, RPCOption};
use openraft::raft::{
AppendEntriesRequest, AppendEntriesResponse, InstallSnapshotRequest, InstallSnapshotResponse,
VoteRequest, VoteResponse,
};
use openraft::{RaftTypeConfig, RPCError};
use tokio::sync::RwLock;
use tonic::transport::Channel;
use crate::types::{TypeConfig, NodeId, Node};
// gRPC 服务定义(需要生成 protobuf 代码)
// 定义在 src/raft/proto/raft.proto
pub mod raft_proto {
tonic::include_proto!("raft");
}
/// Raft gRPC 客户端(类似 C++ 的 brpc)
///
/// **实现方案**:使用 gRPC(tonic),与 C++ 版本的 brpc 一致
pub struct KiwiNetwork {
target: NodeId,
node: Node,
/// gRPC 客户端(延迟初始化)
client: Arc<RwLock<Option<raft_proto::raft_service_client::RaftServiceClient<Channel>>>>,
}
impl KiwiNetwork {
/// 获取或创建 gRPC 客户端
async fn get_client(&self) -> Result<raft_proto::raft_service_client::RaftServiceClient<Channel>, RPCError<NodeId, Node, openraft::error::RaftError<NodeId>>> {
let mut client_guard = self.client.write().await;
if client_guard.is_none() {
// 连接到 gRPC 服务器
let addr = format!("http://{}", self.node.rpc_addr);
let channel = Channel::from_shared(addr.clone())
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?
.connect()
.await
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
*client_guard = Some(raft_proto::raft_service_client::RaftServiceClient::new(channel));
}
Ok(client_guard.as_ref()
.ok_or_else(|| RPCError::Network(openraft::error::NetworkError::new(
&std::io::Error::new(
std::io::ErrorKind::NotConnected,
"gRPC client not initialized",
),
)))?
.clone())
}
}
#[async_trait::async_trait]
impl RaftNetwork<TypeConfig> for KiwiNetwork {
async fn append_entries(
&mut self,
req: AppendEntriesRequest<TypeConfig>,
_option: RPCOption,
) -> Result<AppendEntriesResponse<NodeId>, RPCError<NodeId, Node, openraft::error::RaftError<NodeId>>> {
// 实现方案:使用 gRPC 发送 AppendEntries RPC(参考 C++ 版本的 brpc)
let mut client = self.get_client().await?;
// 序列化请求(openraft 的内部 RPC 类型使用 bincode 序列化)
// 注意:这里序列化的是 openraft 的 AppendEntriesRequest,不是 Binlog
// Binlog 的序列化由 openraft 在日志存储层通过 serde_json 完成
let req_bytes = bincode::serialize(&req)
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
// 构建 gRPC 请求
let grpc_req = tonic::Request::new(raft_proto::AppendEntriesRequest {
data: req_bytes,
});
// 发送 RPC
let grpc_resp = client.append_entries(grpc_req).await
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
// 反序列化响应
let resp_bytes = grpc_resp.into_inner().data;
let resp: AppendEntriesResponse<NodeId> = bincode::deserialize(&resp_bytes)
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
Ok(resp)
}
async fn install_snapshot(
&mut self,
req: InstallSnapshotRequest<TypeConfig>,
_option: RPCOption,
) -> Result<InstallSnapshotResponse<NodeId>, RPCError<NodeId, Node, openraft::error::RaftError<NodeId>>> {
// 实现方案:使用 gRPC 发送 InstallSnapshot RPC
let mut client = self.get_client().await?;
// 序列化请求(openraft 的内部 RPC 类型使用 bincode 序列化)
let req_bytes = bincode::serialize(&req)
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
// 构建 gRPC 请求
let grpc_req = tonic::Request::new(raft_proto::InstallSnapshotRequest {
data: req_bytes,
});
// 发送 RPC
let grpc_resp = client.install_snapshot(grpc_req).await
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
// 反序列化响应
let resp_bytes = grpc_resp.into_inner().data;
let resp: InstallSnapshotResponse<NodeId> = bincode::deserialize(&resp_bytes)
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
Ok(resp)
}
async fn vote(
&mut self,
req: VoteRequest<NodeId>,
_option: RPCOption,
) -> Result<VoteResponse<NodeId>, RPCError<NodeId, Node, openraft::error::RaftError<NodeId>>> {
// 实现方案:使用 gRPC 发送 Vote RPC
let mut client = self.get_client().await?;
// 序列化请求(openraft 的内部 RPC 类型使用 bincode 序列化)
let req_bytes = bincode::serialize(&req)
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
// 构建 gRPC 请求
let grpc_req = tonic::Request::new(raft_proto::VoteRequest {
data: req_bytes,
});
// 发送 RPC
let grpc_resp = client.vote(grpc_req).await
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
// 反序列化响应
let resp_bytes = grpc_resp.into_inner().data;
let resp: VoteResponse<NodeId> = bincode::deserialize(&resp_bytes)
.map_err(|e| RPCError::Network(openraft::error::NetworkError::new(&e)))?;
Ok(resp)
}
}
/// RaftNetworkFactory 实现
///
/// **实现方案**:存储节点地址映射(参考 C++ 版本和 openraft 示例)
pub struct KiwiNetworkFactory {
/// 节点地址映射(推荐方案)
nodes: Arc<RwLock<HashMap<NodeId, Node>>>,
}
impl KiwiNetworkFactory {
pub fn new(nodes: HashMap<NodeId, Node>) -> Self {
Self {
nodes: Arc::new(RwLock::new(nodes)),
}
}
/// 更新节点信息
pub async fn update_node(&self, node_id: NodeId, node: Node) {
let mut nodes = self.nodes.write().await;
nodes.insert(node_id, node);
}
}
#[async_trait::async_trait]
impl RaftNetworkFactory<TypeConfig> for KiwiNetworkFactory {
type Network = KiwiNetwork;
async fn new_client(&mut self, target: NodeId, node: &Node) -> Self::Network {
// 更新节点信息(推荐方案)
{
let mut nodes = self.nodes.write().await;
nodes.insert(target, node.clone());
}
KiwiNetwork {
target,
node: node.clone(),
client: Arc::new(RwLock::new(None)),
}
}
}gRPC Protobuf 定义(src/raft/proto/raft.proto):
syntax = "proto3";
package raft;
// Raft 服务定义
service RaftService {
rpc AppendEntries(AppendEntriesRequest) returns (AppendEntriesResponse);
rpc InstallSnapshot(InstallSnapshotRequest) returns (InstallSnapshotResponse);
rpc Vote(VoteRequest) returns (VoteResponse);
}
// 请求/响应消息(使用 bytes 传输序列化后的数据)
message AppendEntriesRequest {
bytes data = 1;
}
message AppendEntriesResponse {
bytes data = 1;
}
message InstallSnapshotRequest {
bytes data = 1;
}
message InstallSnapshotResponse {
bytes data = 1;
}
message VoteRequest {
bytes data = 1;
}
message VoteResponse {
bytes data = 1;
}依赖配置(src/raft/Cargo.toml):
[dependencies]
tonic = "0.11"
prost = "0.13"
prost-types = "0.13"
bincode = "1.3" # 仅用于 gRPC 网络层序列化 openraft 的内部 RPC 类型
[build-dependencies]
tonic-build = "0.11"构建脚本(src/raft/build.rs):
fn main() -> Result<(), Box<dyn std::error::Error>> {
// 编译 binlog.proto,生成 Rust 代码
// 启用 serde 支持,以满足 openraft 的 Serialize/Deserialize 要求
prost_build::Config::new()
.type_attribute(".kiwi.Binlog", "#[derive(serde::Serialize, serde::Deserialize)]")
.type_attribute(".kiwi.BinlogEntry", "#[derive(serde::Serialize, serde::Deserialize)]")
.type_attribute(".kiwi.OperateType", "#[derive(serde::Serialize, serde::Deserialize)]")
.compile_protos(&["proto/binlog.proto"], &["proto/"])?;
// 编译 gRPC 服务定义
tonic_build::compile_protos("proto/raft.proto")?;
Ok(())
}gRPC 服务端实现(src/raft/src/server.rs):
// src/raft/src/server.rs
use tonic::{Request, Response, Status};
use openraft::Raft;
use crate::network::raft_proto;
use crate::types::TypeConfig;
/// Raft gRPC 服务端实现
pub struct RaftService {
raft: Arc<Raft<TypeConfig>>,
}
#[tonic::async_trait]
impl raft_proto::raft_service_server::RaftService for RaftService {
async fn append_entries(
&self,
request: Request<raft_proto::AppendEntriesRequest>,
) -> Result<Response<raft_proto::AppendEntriesResponse>, Status> {
// 反序列化请求
let req_bytes = request.into_inner().data;
let req: openraft::raft::AppendEntriesRequest<TypeConfig> = bincode::deserialize(&req_bytes)
.map_err(|e| Status::invalid_argument(format!("Failed to deserialize request: {}", e)))?;
// 调用 Raft
let resp = self.raft.append_entries(req).await
.map_err(|e| Status::internal(format!("Raft error: {}", e)))?;
// 序列化响应
let resp_bytes = bincode::serialize(&resp)
.map_err(|e| Status::internal(format!("Failed to serialize response: {}", e)))?;
Ok(Response::new(raft_proto::AppendEntriesResponse {
data: resp_bytes,
}))
}
async fn install_snapshot(
&self,
request: Request<raft_proto::InstallSnapshotRequest>,
) -> Result<Response<raft_proto::InstallSnapshotResponse>, Status> {
// 反序列化请求
let req_bytes = request.into_inner().data;
let req: openraft::raft::InstallSnapshotRequest<TypeConfig> = bincode::deserialize(&req_bytes)
.map_err(|e| Status::invalid_argument(format!("Failed to deserialize request: {}", e)))?;
// 调用 Raft
let resp = self.raft.install_snapshot(req).await
.map_err(|e| Status::internal(format!("Raft error: {}", e)))?;
// 序列化响应
let resp_bytes = bincode::serialize(&resp)
.map_err(|e| Status::internal(format!("Failed to serialize response: {}", e)))?;
Ok(Response::new(raft_proto::InstallSnapshotResponse {
data: resp_bytes,
}))
}
async fn vote(
&self,
request: Request<raft_proto::VoteRequest>,
) -> Result<Response<raft_proto::VoteResponse>, Status> {
// 反序列化请求
let req_bytes = request.into_inner().data;
let req: openraft::raft::VoteRequest<u64> = bincode::deserialize(&req_bytes)
.map_err(|e| Status::invalid_argument(format!("Failed to deserialize request: {}", e)))?;
// 调用 Raft
let resp = self.raft.vote(req).await
.map_err(|e| Status::internal(format!("Raft error: {}", e)))?;
// 序列化响应
let resp_bytes = bincode::serialize(&resp)
.map_err(|e| Status::internal(format!("Failed to serialize response: {}", e)))?;
Ok(Response::new(raft_proto::VoteResponse {
data: resp_bytes,
}))
}
}文件:src/raft/src/node.rs
// src/raft/src/node.rs
use std::sync::Arc;
use openraft::{Config, Raft};
use tokio::sync::RwLock;
use crate::types::{TypeConfig, NodeId, Node, Request, Response};
use crate::log_store::KiwiLogStore;
use crate::state_machine::KiwiStateMachine;
use crate::network::{KiwiNetworkFactory, KiwiNetwork};
use crate::error::RaftResult;
/// RaftNode 接口(类似 C++ 的 Raft 类)
#[async_trait::async_trait]
pub trait RaftNodeInterface: Send + Sync {
/// 检查是否是 Leader
async fn is_leader(&self) -> bool;
/// 获取 Leader 地址
async fn get_leader_address(&self) -> Option<String>;
/// 追加日志到 Raft(类似 C++ 的 AppendLog)
async fn append_log(&self, request: Request) -> RaftResult<Response>;
/// 检查是否已初始化
fn is_initialized(&self) -> bool;
}
/// RaftNode 封装(类似 C++ 的 Raft 类)
pub struct RaftNode {
raft: Arc<Raft<TypeConfig>>,
node_id: NodeId,
initialized: Arc<RwLock<bool>>,
}
impl RaftNode {
/// 创建新的 RaftNode
pub async fn new(
node_id: NodeId,
config: Arc<Config>,
log_store: KiwiLogStore,
state_machine: KiwiStateMachine,
network_factory: KiwiNetworkFactory,
node_configs: HashMap<NodeId, Node>,
) -> RaftResult<Self> {
let raft = Raft::new(node_id, config, network_factory, log_store, state_machine)
.await
.map_err(|e| crate::error::RaftError::from(e))?;
Ok(Self {
raft: Arc::new(raft),
node_id,
initialized: Arc::new(RwLock::new(false)),
node_configs: Arc::new(RwLock::new(node_configs)),
})
}
/// 初始化集群(单节点)
pub async fn initialize(&self) -> RaftResult<()> {
self.raft.initialize(vec![self.node_id]).await
.map_err(|e| crate::error::RaftError::from(e))?;
let mut init = self.initialized.write().await;
*init = true;
Ok(())
}
}
#[async_trait::async_trait]
impl RaftNodeInterface for RaftNode {
async fn is_leader(&self) -> bool {
let metrics = self.raft.metrics().borrow().clone();
metrics.current_leader_id == Some(self.node_id)
}
async fn get_leader_address(&self) -> Option<String> {
let metrics = self.raft.metrics().borrow().clone();
if let Some(leader_id) = metrics.current_leader_id {
// 从 node_configs 中获取 Leader 的地址(生产级实现,对应 C++ 的 leader_id().addr)
let node_configs = self.node_configs.read().await;
if let Some(leader_node) = node_configs.get(&leader_id) {
// 返回 RPC 地址(对应 C++ 的 endpoint2str)
Some(leader_node.rpc_addr.clone())
} else {
None
}
} else {
None
}
}
async fn append_log(&self, request: Request) -> RaftResult<Response> {
// 检查是否是 Leader
if !self.is_leader().await {
return Err(crate::error::RaftError::NotLeader {
leader_id: self.get_leader_address().await,
});
}
// 直接提交 Binlog 到 Raft
// openraft 会使用 serde(serde_json)序列化整个 Entry(包括 EntryPayload::Normal(request))
// 由于 Binlog 通过 prost-build 的 serde 特性实现了 Serialize/Deserialize,
// openraft 可以直接序列化它,无需手动转换
let result = self.raft.client_write(request).await
.map_err(|e| crate::error::RaftError::from(e))?;
// 从状态机获取响应
// openraft 的 client_write 返回的 ClientWriteResponse 包含 data 字段,
// 即状态机返回的 Response
Ok(result.data)
}
fn is_initialized(&self) -> bool {
// 检查 Raft 节点是否已初始化(生产级实现,对应 C++ 的 IsInitialized)
// 对应 C++ 的 node_ != nullptr && server_ != nullptr
// openraft 提供了 is_initialized 方法,但需要在异步上下文中调用
// 这里使用缓存的 initialized 标志
self.initialized.load(std::sync::atomic::Ordering::Acquire)
}
}文件:src/storage/src/batch.rs
// src/storage/src/batch.rs
use crate::error::Result;
use crate::storage::Redis;
use crate::raft::RaftNodeInterface;
use std::sync::Arc;
use crate::raft::types::{Binlog, BinlogEntry, OperateType};
/// ColumnFamily 索引
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ColumnFamilyIndex {
MetaCF = 0,
HashesDataCF = 1,
SetsDataCF = 2,
ListsDataCF = 3,
ZsetsDataCF = 4,
ZsetsScoreCF = 5,
}
/// Batch trait(类似 C++ 的 Batch)
#[async_trait::async_trait]
pub trait Batch: Send + Sync {
fn put(&mut self, cf: ColumnFamilyIndex, key: &[u8], value: &[u8]);
fn delete(&mut self, cf: ColumnFamilyIndex, key: &[u8]);
async fn commit(&mut self) -> Result<()>;
fn count(&self) -> usize;
}
/// RocksBatch:直接写入 RocksDB(单机模式,与 C++ 版本一致)
/// 优化:在 Put/Delete 时直接操作 WriteBatch,而不是先收集到 Vec
pub struct RocksBatch {
batch: rocksdb::WriteBatch, // 直接使用 WriteBatch(对应 C++ 的 batch_)
db: Arc<rocksdb::DB>,
write_options: rocksdb::WriteOptions,
cf_handles: Arc<Vec<rocksdb::ColumnFamilyHandle>>,
cnt: usize,
}
impl RocksBatch {
pub fn new(
db: Arc<rocksdb::DB>,
write_options: rocksdb::WriteOptions,
cf_handles: Arc<Vec<rocksdb::ColumnFamilyHandle>>,
) -> Self {
Self {
batch: rocksdb::WriteBatch::default(),
db,
write_options,
cf_handles,
cnt: 0,
}
}
}
#[async_trait::async_trait]
impl Batch for RocksBatch {
fn put(&mut self, cf: ColumnFamilyIndex, key: &[u8], value: &[u8]) {
// 直接操作 WriteBatch(对应 C++ 的 batch_.Put)
if let Some(cf_handle) = self.cf_handles.get(cf as usize) {
self.batch.put_cf(cf_handle, key, value);
self.cnt += 1;
}
}
fn delete(&mut self, cf: ColumnFamilyIndex, key: &[u8]) {
// 直接操作 WriteBatch(对应 C++ 的 batch_.Delete)
if let Some(cf_handle) = self.cf_handles.get(cf as usize) {
self.batch.delete_cf(cf_handle, key);
self.cnt += 1;
}
}
async fn commit(&mut self) -> Result<()> {
// 直接写入 RocksDB(对应 C++ 的 db_->Write(options_, &batch_))
self.db.write_opt(&self.batch, &self.write_options)
.map_err(|e| crate::error::Error::from(e))?;
Ok(())
}
fn count(&self) -> usize {
self.cnt
}
}
/// RaftBatch:通过 Raft 复制(Raft 模式,操作级日志,与 C++ 版本一致)
pub struct RaftBatch {
binlog: Binlog, // 类似 C++ 的 BinlogBatch::binlog_
raft_node: Arc<dyn RaftNodeInterface>,
instance_id: u32,
timeout_seconds: u32, // 超时时间(对应 C++ 的 seconds_)
cnt: usize,
}
impl RaftBatch {
pub fn new(
raft_node: Arc<dyn RaftNodeInterface>,
instance_id: u32,
timeout_seconds: u32,
) -> Self {
Self {
binlog: Binlog {
db_id: 0,
slot_idx: instance_id,
entries: Vec::new(),
},
raft_node,
instance_id,
timeout_seconds,
cnt: 0,
}
}
}
#[async_trait::async_trait]
impl Batch for RaftBatch {
fn put(&mut self, cf: ColumnFamilyIndex, key: &[u8], value: &[u8]) {
// 添加到 Binlog(类似 C++ 的 BinlogBatch::Put)
// 注意:protobuf 生成的类型使用 i32 表示枚举
self.binlog.entries.push(BinlogEntry {
cf_idx: cf as u32,
op_type: OperateType::KPut as i32,
key: key.to_vec(),
value: Some(value.to_vec()),
});
self.cnt += 1;
}
fn delete(&mut self, cf: ColumnFamilyIndex, key: &[u8]) {
// 添加到 Binlog(类似 C++ 的 BinlogBatch::Delete)
self.binlog.entries.push(BinlogEntry {
cf_idx: cf as u32,
op_type: OperateType::KDelete as i32,
key: key.to_vec(),
value: None,
});
self.cnt += 1;
}
async fn commit(&mut self) -> Result<()> {
if self.binlog.entries.is_empty() {
return Ok(());
}
// 提交 Binlog 到 Raft(类似 C++ 的 BinlogBatch::Commit)
// 添加超时机制(对应 C++ 的 future.wait_for(std::chrono::seconds(seconds_)))
match tokio::time::timeout(
std::time::Duration::from_secs(self.timeout_seconds as u64),
self.raft_node.append_log(self.binlog.clone())
).await {
Ok(Ok(())) => Ok(()),
Ok(Err(e)) => Err(e),
Err(_) => {
// 超时(对应 C++ 的 Status::Incomplete("Wait for write timeout"))
Err(crate::error::Error::Timeout(
"Wait for write timeout".to_string()
))
}
}
}
fn count(&self) -> usize {
self.cnt
}
}
/// Batch 工厂函数(类似 C++ 的 Batch::CreateBatch)
/// 从 Redis 实例获取所有信息,与 C++ 版本一致
pub fn create_batch(redis: &Redis) -> Box<dyn Batch> {
// 检查是否有 RaftNode(对应 C++ 的 redis->GetAppendLogFunction())
if let Some(raft_node) = redis.get_raft_node() {
// 创建 RaftBatch(对应 C++ 的 BinlogBatch)
Box::new(RaftBatch::new(
raft_node,
redis.get_index() as u32,
redis.get_raft_timeout(),
))
} else {
// 创建 RocksBatch(对应 C++ 的 RocksBatch)
Box::new(RocksBatch::new(
redis.get_db(),
redis.get_write_options(),
redis.get_column_family_handles(),
))
}
}说明:
- ✅ 与 C++ 版本一致:
create_batch接受&Redis参数,从Redis实例获取所有信息 - ✅ 超时机制:
RaftBatch::commit()使用tokio::time::timeout,对应 C++ 的future.wait_for - ✅ 性能优化:
RocksBatch在Put/Delete时直接操作WriteBatch,而不是先收集到Vec - ✅ 架构一致:Batch 在 Redis 层创建和使用,与 C++ 版本完全一致
文件:src/raft/src/error.rs
// src/raft/src/error.rs
use crate::types::NodeId;
use thiserror::Error;
/// Raft 错误类型
#[derive(Error, Debug)]
pub enum RaftError {
#[error("Not leader, leader: {:?}", leader_id)]
NotLeader { leader_id: Option<String> },
#[error("Serialization error: {0}")]
Serialization(#[from] bincode::Error),
#[error("Storage error: {0}")]
Storage(#[from] openraft::StorageError<NodeId>),
#[error("Raft error: {0}")]
Raft(#[from] openraft::error::RaftError<NodeId>),
#[error("IO error: {0}")]
IO(#[from] std::io::Error),
}
pub type RaftResult<T> = Result<T, RaftError>;文件:src/raft/src/mod.rs
// src/raft/src/mod.rs
pub mod types;
pub mod log_index;
pub mod state_machine;
pub mod flush_manager;
pub mod snapshot;
pub mod log_store;
pub mod network;
pub mod node;
pub mod error;
pub use types::*;
pub use log_index::*;
pub use state_machine::KiwiStateMachine;
pub use flush_manager::FlushEventListener;
pub use snapshot::SnapshotManager;
pub use log_store::KiwiLogStore;
pub use network::{KiwiNetworkFactory, KiwiNetwork};
pub use node::{RaftNode, RaftNodeInterface};
pub use error::{RaftError, RaftResult};文件:src/storage/src/storage.rs(修改部分)
说明:
- ✅ Storage 层提供 Checkpoint 方法:
create_checkpoint和load_checkpoint由 Storage 层实现,管理多个 Redis 实例 - ✅ Batch 方法实现:
batch_put_cf和batch_delete_cf根据 key 的 slot_id 路由到对应的 Redis 实例 - ✅ 与 C++ 版本一致:Storage 层负责管理多个 RocksDB 实例的 checkpoint 操作
- ✅ LoadCheckpoint 调用时机:
load_checkpoint在 Storage 打开之前调用,不依赖self.insts,接受db_instance_num参数 - ✅ LoadCheckpointInternal 逻辑:每个实例的备份和恢复逻辑与 C++ 版本一致(重命名为 .tmp,复制文件,删除备份)
文件:src/storage/src/storage_impl.rs(修改部分)
// src/storage/src/storage_impl.rs(修改部分)
impl Storage {
/// SET 命令(路由层,与 C++ 版本一致)
/// 只负责路由到对应的 Redis 实例,不创建 Batch
pub async fn set(&self, key: &[u8], value: &[u8]) -> Result<()> {
// 路由到对应的 Redis 实例(对应 C++ 的 GetDBInstance)
let slot_id = key_to_slot_id(key);
let instance_id = self.slot_indexer.get_instance_id(slot_id);
let redis = &self.insts[instance_id];
// 调用 Redis 层的方法(对应 C++ 的 inst->Set(key, value))
redis.set(key, value).await
}
/// GET 命令(路由层,直接读取,不经过 Raft)
pub async fn get(&self, key: &[u8]) -> Result<String> {
// 路由到对应的 Redis 实例
let slot_id = key_to_slot_id(key);
let instance_id = self.slot_indexer.get_instance_id(slot_id);
self.insts[instance_id].get(key).await
}文件:src/storage/src/redis_strings.rs(新增/修改部分)
// src/storage/src/redis_strings.rs(新增/修改部分)
use crate::batch::{Batch, create_batch, ColumnFamilyIndex};
use crate::storage::Redis;
impl Redis {
/// SET 命令(命令实现层,与 C++ 版本一致)
/// 在这里创建和使用 Batch,对应 C++ 的 Redis::Set()
pub async fn set(&self, key: &[u8], value: &[u8]) -> Result<()> {
// 创建 Batch(对应 C++ 的 Batch::CreateBatch(this))
let mut batch = create_batch(self);
// 编码 key/value
let base_key = BaseKey::new(key);
let strings_value = StringsValue::new(value);
// 添加到 Batch(操作级,不是命令级)
// 对应 C++ 的 batch->Put(kMetaCF, base_key.Encode(), strings_value.Encode())
batch.put(
ColumnFamilyIndex::MetaCF,
&base_key.encode(),
&strings_value.encode(),
);
// 提交(如果是 RaftBatch,会通过 Raft 复制 Binlog)
// 对应 C++ 的 batch->Commit()
batch.commit().await?;
Ok(())
}
/// 获取 RaftNode(用于 Batch 创建)
pub fn get_raft_node(&self) -> Option<Arc<dyn RaftNodeInterface>> {
self.raft_node.clone()
}
/// 获取 Raft 超时时间(用于 Batch 创建)
pub fn get_raft_timeout(&self) -> u32 {
self.raft_timeout_seconds
}
/// 获取实例索引(用于 Batch 创建)
pub fn get_index(&self) -> usize {
self.index
}
/// 获取 DB(用于 Batch 创建)
pub fn get_db(&self) -> Arc<rocksdb::DB> {
self.db.clone()
}
/// 获取 WriteOptions(用于 Batch 创建)
pub fn get_write_options(&self) -> rocksdb::WriteOptions {
self.default_write_options.clone()
}
/// 获取 ColumnFamily handles(用于 Batch 创建)
pub fn get_column_family_handles(&self) -> Arc<Vec<rocksdb::ColumnFamilyHandle>> {
self.cf_handles.clone()
}
/// 批量 Put(用于状态机应用 Binlog)
/// 根据 key 的 slot_id 路由到对应的 Redis 实例
pub async fn batch_put_cf(
&self,
cf: ColumnFamilyIndex,
key: &[u8],
value: &[u8],
) -> Result<()> {
// 根据 key 获取对应的 Redis 实例(对应 C++ 的 GetDBInstance)
let slot_id = crate::slot_indexer::key_to_slot_id(key);
let instance_id = self.slot_indexer.get_instance_id(slot_id);
let instance = &self.insts[instance_id];
// 获取对应的 ColumnFamily handle
let cf_handle = instance.get_cf_handle(cf)?;
// 直接写入 RocksDB(对应 C++ 的 db->Put)
instance.db.as_ref()
.ok_or_else(|| Error::OptionNone {
message: "db is not initialized".to_string(),
})?
.put_cf(cf_handle, key, value)?;
Ok(())
}
/// 批量 Delete(用于状态机应用 Binlog)
/// 根据 key 的 slot_id 路由到对应的 Redis 实例
pub async fn batch_delete_cf(
&self,
cf: ColumnFamilyIndex,
key: &[u8],
) -> Result<()> {
// 根据 key 获取对应的 Redis 实例(对应 C++ 的 GetDBInstance)
let slot_id = crate::slot_indexer::key_to_slot_id(key);
let instance_id = self.slot_indexer.get_instance_id(slot_id);
let instance = &self.insts[instance_id];
// 获取对应的 ColumnFamily handle
let cf_handle = instance.get_cf_handle(cf)?;
// 直接删除(对应 C++ 的 db->Delete)
instance.db.as_ref()
.ok_or_else(|| Error::OptionNone {
message: "db is not initialized".to_string(),
})?
.delete_cf(cf_handle, key)?;
Ok(())
}
/// 创建 Checkpoint(对应 C++ 的 Storage::CreateCheckpoint)
/// 为每个 Redis 实例创建 checkpoint,返回异步任务句柄
pub async fn create_checkpoint(
&self,
checkpoint_path: &std::path::Path,
) -> Result<Vec<tokio::task::JoinHandle<Result<()>>>, Box<dyn std::error::Error>> {
use rocksdb::checkpoint::Checkpoint;
let mut handles = Vec::new();
for (i, inst) in self.insts.iter().enumerate() {
let checkpoint_sub_path = checkpoint_path.join(i.to_string());
let db_path = inst.get_db_path()?; // 需要从 Redis 获取路径
// 异步执行(对应 C++ 的 std::async(std::launch::async, ...))
let handle = tokio::task::spawn_blocking(move || -> Result<()> {
// 创建 Checkpoint(对应 C++ 的 Storage::CreateCheckpointInternal)
let db = rocksdb::DB::open_for_read_only(
&rocksdb::Options::default(),
&db_path,
false,
)?;
// 创建临时目录(对应 C++ 的 tmp_dir)
let tmp_dir = checkpoint_sub_path.with_extension("tmp");
if tmp_dir.exists() {
std::fs::remove_dir_all(&tmp_dir)?;
}
let checkpoint = Checkpoint::new(&db)?;
checkpoint.create_checkpoint(&tmp_dir)?;
// 删除源目录(如果存在)
if checkpoint_sub_path.exists() {
std::fs::remove_dir_all(&checkpoint_sub_path)?;
}
// 重命名临时目录为源目录(对应 C++ 的 RenameFile)
std::fs::rename(&tmp_dir, &checkpoint_sub_path)?;
Ok(())
});
handles.push(handle);
}
Ok(handles)
}
/// 从 Checkpoint 加载数据库(对应 C++ 的 Storage::LoadCheckpoint)
/// 为每个 Redis 实例从 checkpoint 加载数据,返回异步任务句柄
///
/// **注意**:这是一个静态方法或需要在 Storage 打开之前调用
/// 因为 C++ 版本中,LoadCheckpoint 是在 Storage::Open() 之前调用的
/// 此时 Storage 的 insts_ 还未初始化,所以不能依赖 self.insts
pub async fn load_checkpoint(
checkpoint_path: &std::path::Path,
db_path: &std::path::Path,
db_instance_num: usize,
) -> Result<Vec<tokio::task::JoinHandle<Result<()>>>, Box<dyn std::error::Error>> {
let mut handles = Vec::new();
// 使用传入的 db_instance_num,而不是 self.insts.len()
// 因为此时 Storage 还没打开,insts 为空
for i in 0..db_instance_num {
let checkpoint_sub_path = checkpoint_path.join(i.to_string());
let db_sub_path = db_path.join(i.to_string());
// 异步执行(对应 C++ 的 std::async(std::launch::async, ...))
let handle = tokio::task::spawn_blocking(move || -> Result<()> {
// 从 Checkpoint 复制文件(对应 C++ 的 Storage::LoadCheckpointInternal)
// 1. 重命名原数据库为 .tmp(备份)
let tmp_path = db_sub_path.with_extension("tmp");
if db_sub_path.exists() {
std::fs::rename(&db_sub_path, &tmp_path)?;
}
// 2. 创建新的数据库目录
std::fs::create_dir_all(&db_sub_path)?;
// 3. 从 Checkpoint 复制文件(对应 C++ 的 RecursiveLinkAndCopy)
if let Err(e) = copy_dir_all(&checkpoint_sub_path, &db_sub_path) {
// 回滚:恢复原数据库
if tmp_path.exists() {
let _ = std::fs::remove_dir_all(&db_sub_path);
let _ = std::fs::rename(&tmp_path, &db_sub_path);
}
return Err(e.into());
}
// 4. 删除备份目录(对应 C++ 的 DestroyDB)
if tmp_path.exists() {
let _ = std::fs::remove_dir_all(&tmp_path);
}
Ok(())
});
handles.push(handle);
}
Ok(handles)
}
}
// 辅助函数:递归复制目录(对应 C++ 的 RecursiveLinkAndCopy)
fn copy_dir_all(src: impl AsRef<std::path::Path>, dst: impl AsRef<std::path::Path>) -> std::io::Result<()> {
std::fs::create_dir_all(&dst)?;
for entry in std::fs::read_dir(src)? {
let entry = entry?;
let ty = entry.file_type()?;
if ty.is_dir() {
copy_dir_all(entry.path(), dst.as_ref().join(entry.file_name()))?;
} else {
std::fs::copy(entry.path(), dst.as_ref().join(entry.file_name()))?;
}
}
Ok(())
}文件:src/cmd/src/lib.rs(修改部分)
// src/cmd/src/lib.rs(修改部分)
use crate::raft::RaftNodeInterface;
use std::sync::Arc;
pub trait Cmd: Send + Sync {
// ... 现有方法 ...
fn execute(&self, client: &Client, storage: Arc<Storage>, raft_node: Option<Arc<dyn RaftNodeInterface>>) {
debug!("execute command: {:?}", client.cmd_name());
// 【关键】Leader 检查(类似 C++ 的 BaseCmd::Execute)
if let Some(raft) = &raft_node {
let rt = tokio::runtime::Handle::current();
let is_leader = rt.block_on(raft.is_leader());
if !is_leader {
// 获取 Leader 地址
let leader_addr = rt.block_on(raft.get_leader_address());
if let Some(addr) = leader_addr {
// 计算 hash slot
let hash_slot = crate::storage::slot_indexer::key_to_slot_id(client.key());
client.set_reply(RespData::Error(
format!("MOVED {} {}", hash_slot, addr).into()
));
} else {
client.set_reply(RespData::Error("CLUSTERDOWN No raft leader".into()));
}
return;
}
}
if !self.check_arg(client.argv().len()) {
client.set_reply(RespData::Error(
format!(
"ERR wrong number of arguments for '{}' command",
String::from_utf8_lossy(client.cmd_name().as_slice()),
)
.into(),
));
return;
}
if self.do_initial(client) {
self.do_cmd(client, storage, raft_node);
}
}
fn do_cmd(&self, client: &Client, storage: Arc<Storage>, raft_node: Option<Arc<dyn RaftNodeInterface>>);
}文件:src/raft/src/lib.rs
// src/raft/src/lib.rs
use std::path::Path;
use std::sync::Arc;
use openraft::Config;
use storage::storage::Storage;
use rocksdb::DB;
use crate::log_store::KiwiLogStore;
use crate::state_machine::KiwiStateMachine;
use crate::network::KiwiNetworkFactory;
use crate::node::RaftNode;
use crate::flush_manager::FlushEventListener;
/// 创建 Raft 存储(类似 openraft 示例的 new_storage)
pub async fn new_raft_storage<P: AsRef<Path>>(
db_path: P,
db_id: u32,
storage: Arc<Storage>,
db: Arc<DB>,
cf_handles: Arc<Vec<rocksdb::ColumnFamilyHandle>>,
cf_count: usize,
) -> Result<(KiwiLogStore, KiwiStateMachine), crate::error::RaftError> {
// 创建日志存储(文件系统)
let log_store = KiwiLogStore::new(&db_path, db_id)?;
// 创建状态机
let state_machine = KiwiStateMachine::new(storage, db, cf_handles.clone(), cf_count);
Ok((log_store, state_machine))
}
/// 创建 RaftNode
pub async fn create_raft_node(
node_id: crate::types::NodeId,
config: Arc<Config>,
log_store: KiwiLogStore,
state_machine: KiwiStateMachine,
network_factory: KiwiNetworkFactory,
node_configs: HashMap<NodeId, Node>,
) -> Result<RaftNode, crate::error::RaftError> {
RaftNode::new(node_id, config, log_store, state_machine, network_factory, node_configs).await
}
/// 默认 Raft 配置
pub fn default_raft_config() -> Config {
Config {
heartbeat_interval: 250,
election_timeout_min: 299,
election_timeout_max: 500,
..Default::default()
}
.validate()
.map_err(|e| {
tracing::error!("Invalid Raft config: {:?}", e);
std::io::Error::new(
std::io::ErrorKind::InvalidInput,
format!("Invalid Raft config: {}", e),
)
})?
}客户端请求:SET key value
│
├─→ NetworkHandler::handle_connection()
│ │
│ └─→ CmdTable::execute()
│ │
│ └─→ SetCmd::execute()
│ │
│ ├─→ 【Leader 检查】
│ │ ├─→ raft_node.is_leader()?
│ │ │ ├─→ 否 → 返回 MOVED 错误
│ │ │ └─→ 是 → 继续
│ │ │
│ │ └─→ SetCmd::do_cmd()
│ │ │
│ │ └─→ Storage::set() // Storage 层:路由层
│ │ │
│ │ └─→ Redis::set() // Redis 层:命令实现层
│ │ │
│ │ ├─→ create_batch(self) // 【关键】在 Redis 层创建 Batch
│ │ │ ├─→ 如果有 RaftNode → RaftBatch
│ │ │ └─→ 如果没有 → RocksBatch
│ │ │
│ │ ├─→ batch.put(CF, key, value)
│ │ │ └─→ binlog.entries.push(BinlogEntry { Put, ... })
│ │ │
│ │ └─→ batch.commit()
│ │ │
│ │ └─→ RaftBatch::commit()
│ │ │
│ │ ├─→ 构建 Request::Binlog
│ │ │
│ │ ├─→ RaftNode::append_log()(带超时)
│ │ │ │
│ │ │ ├─→ 检查:is_leader()?
│ │ │ │
│ │ │ └─→ raft.client_write()
│ │ │ │
│ │ │ └─→ openraft 内部
│ │ │ │
│ │ │ ├─→ KiwiLogStore::append()
│ │ │ │ └─→ 写入文件系统日志文件
│ │ │ │ └─→ db_path/db_id/_praft/log/{index}.log
│ │ │ │
│ │ │ ├─→ 复制到 Follower
│ │ │ │
│ │ │ ├─→ 等待大多数确认
│ │ │ │
│ │ │ └─→ 提交(committed)
│ │ │ │
│ │ │ └─→ KiwiStateMachine::apply()
│ │ │ │
│ │ │ ├─→ 检查 applied_index
│ │ │ │
│ │ │ ├─→ 匹配 Request::Binlog
│ │ │ │
│ │ │ └─→ apply_binlog()
│ │ │ │
│ │ │ ├─→ 遍历 entries
│ │ │ │
│ │ │ ├─→ 检查 IsApplied(cf_idx, log_idx)
│ │ │ │
│ │ │ ├─→ 直接执行 Put/Delete
│ │ │ │ └─→ db.put_cf() / db.delete_cf()
│ │ │ │
│ │ │ └─→ 更新 CF applied_index
│ │ │ │
│ │ └─→ 如果超时,返回 Timeout 错误
│ │ │
└─────────────────────────────────────────────────────────────────────────────────────────┘
│
└─→ 返回响应给客户端
客户端请求:GET key
│
├─→ NetworkHandler::handle_connection()
│ │
│ └─→ CmdTable::execute()
│ │
│ └─→ GetCmd::execute()
│ │
│ ├─→ 【Leader 检查】
│ │ ├─→ raft_node.is_leader()?
│ │ │ ├─→ 否 → 返回 MOVED 错误
│ │ │ └─→ 是 → 继续(Leader Lease 读)
│ │ │
│ │ └─→ GetCmd::do_cmd()
│ │ │
│ │ └─→ Storage::get()
│ │ │
│ │ └─→ 直接从 RocksDB 读取
│ │ │
│ │ └─→ 不经过 Raft!
│ │ │
└───────────────────────────────────────────────┘
│
└─→ 返回结果给客户端
| 组件 | C++ 版本 | Rust 版本 |
|---|---|---|
| 日志类型 | 操作级(Binlog) | 操作级(Binlog) |
| Request 类型 | Binlog |
Binlog |
| Binlog 定义 | Protobuf (binlog.proto) |
Protobuf (binlog.proto) |
| 序列化方式 | Protobuf(直接序列化 Binlog) | Serde(serde_json 序列化整个 Entry) |
| BinlogEntry | cf_idx, op_type, key, value |
cf_idx, op_type, key, value |
| 网络协议 | brpc(gRPC) | gRPC(tonic) |
| 状态机应用 | 直接执行 Put/Delete | 直接执行 Put/Delete |
| 索引维护 | 每个 CF 一个索引 | 每个 CF 一个索引 |
| 幂等性检查 | IsApplied(cf_idx, log_idx) |
is_applied(cf_idx, log_idx) |
| Batch 抽象 | Batch::CreateBatch(Redis*) |
create_batch(&Redis) |
| Batch 创建位置 | Redis 层(Redis::Set()) |
Redis 层(Redis::set()) |
| Storage 层职责 | 路由层(Storage::Set()) |
路由层(Storage::set()) |
| Batch 超时机制 | future.wait_for(seconds_) |
tokio::time::timeout() |
| Leader 检查 | BaseCmd::Execute() |
Cmd::execute() |
| 日志追加 | Raft::AppendLog() |
RaftNode::append_log() |
| 日志应用 | Raft::on_apply() |
KiwiStateMachine::apply() |
| 数据应用 | Storage::OnBinlogWrite() |
KiwiStateMachine::apply_binlog() |
| Flush 管理 | LogIndexAndSequenceCollectorPurger |
FlushEventListener |
| 快照管理 | GetSmallestFlushedLogIndex() |
get_last_flush_index() |
| 日志存储 | 文件系统(_praft/log) |
文件系统(_praft/log) |
| 日志组织方式 | braft SegmentLogStorage(段模式) | 段模式(Segment) |
| 文件大小限制 | braft 内部管理 | 64MB(可配置) |
| 段元数据 | braft 内部管理 | segments.json |
| 元数据存储 | 文件系统(_praft/raft_meta) |
文件系统(_praft/raft_meta) |
| 数据存储 | RocksDB(db_path/db_id/) |
RocksDB(db_path/db_id/) |
| DB 结构 | DB 类包含 storage_mutex_ 和 storage_ |
DB 结构体包含 storage_mutex 和 storage |
| 快照传输 | brpc RPC 分块传输(InstallSnapshot) |
gRPC 分块传输(install_snapshot) |
| 快照存储 | 文件系统(Checkpoint 目录) | 文件系统(Checkpoint 目录) |
| 快照创建 | CreateCheckpoint()(共享锁) |
create_checkpoint()(共享锁) |
| 快照加载 | LoadDBFromCheckpoint()(独占锁) |
load_db_from_checkpoint()(独占锁) |
| SnapshotData 类型 | 文件路径(local://) |
tokio::fs::File(文件句柄) |
本设计文档提供了完整的 Raft 集成方案,包括:
-
✅ 所有必需的 Trait 实现:
- TypeConfig(使用
declare_raft_types!) - RaftLogStorage(9 个方法)
- RaftLogReader(1 个方法)
- RaftStateMachine(6 个方法)
- RaftSnapshotBuilder(1 个方法)
- TypeConfig(使用
-
✅ 操作级日志:使用 Binlog,与 C++ 版本一致
-
✅ 每个 CF 一个 Index:维护 applied_index 和 flushed_index
-
✅ Flush 管理:EventListener 监听 Flush 事件,自动触发快照
-
✅ 快照管理:基于最小 flushed_index 创建快照
-
✅ Batch 抽象层:统一 Raft 和单机模式,与 C++ 版本完全一致
- Batch 在 Redis 层创建和使用(
Redis::set()),Storage 层只负责路由 -
create_batch(&Redis)从 Redis 实例获取所有信息,对应 C++ 的Batch::CreateBatch(Redis*) -
RaftBatch::commit()使用tokio::time::timeout实现超时机制,对应 C++ 的future.wait_for -
RocksBatch在Put/Delete时直接操作WriteBatch,性能优化与 C++ 版本一致
- Batch 在 Redis 层创建和使用(
-
✅ Leader 检查:读写命令都需要在 Leader 执行
-
✅ 文件系统日志存储:与 C++ 版本一致,使用文件系统存储 Raft 日志和元数据
-
✅ 段模式日志组织:与 C++ 版本的 braft SegmentLogStorage 一致,使用段(Segment)模式组织日志文件,一个文件包含多个连续的 log_index,文件大小限制为 64MB(可配置)
-
✅ 段元数据管理:使用
segments.json记录每个段的起始和结束 log_index,用于快速查找,与 braft 的内部段管理机制一致 -
✅ Protobuf Binlog:与 C++ 版本一致,使用相同的
binlog.proto定义,通过prost生成 Rust 代码,并通过prost-build的serde特性实现Serialize/Deserialize,满足 openraft 的序列化要求 -
✅ gRPC 网络层:与 C++ 版本的 brpc 一致,使用 tonic 实现 gRPC 服务
-
✅ 序列化机制:Binlog 通过 serde(serde_json)序列化,无需手动转换为 bytes;bincode 仅用于 gRPC 网络层序列化 openraft 的内部 RPC 类型
-
✅ DB 包装结构体:与 C++ 版本一致,
storage_mutex放在DB结构体中,CreateCheckpoint使用共享锁,LoadDBFromCheckpoint使用独占锁 -
✅ 快照分块传输:与 C++ 版本一致,使用 gRPC 的
install_snapshotRPC 分块传输快照数据,支持大文件传输 -
✅ 生产级错误处理:所有操作都使用
Result类型,不使用unwrap()或expect(),确保生产级可靠性 -
✅ 函数命名:
create_checkpoint和load_db_from_checkpoint与 C++ 版本一致,便于理解和维护 -
✅ 架构层次一致:Storage 层作为路由层,Redis 层作为命令实现层,Batch 在 Redis 层创建和使用,与 C++ 版本完全一致
该实现与 C++ 版本的设计理念完全一致,包括存储架构(文件系统日志 + RocksDB 数据)和数据结构(Protobuf Binlog),可以直接用于生产环境。