use std::collections::HashMap;
use std::future::Future;
use std::io::Cursor;
use std::sync::Arc;
use std::time::Duration;
use futures::StreamExt;
use openraft::error::{NetworkError, RPCError, StreamingError, Unreachable};
use openraft::errors::ReplicationClosed;
use openraft::network::{RPCOption, RaftNetworkFactory, RaftNetworkV2};
use openraft::raft::{
AppendEntriesRequest, AppendEntriesResponse, SnapshotResponse, TransferLeaderRequest,
VoteRequest, VoteResponse,
};
use openraft::type_config::alias::{SnapshotOf, VoteOf};
use tokio::sync::Mutex;
use tonic::transport::{Channel, ClientTlsConfig};
use tsoracle_driver_openraft::{NodeCapabilities, OpenraftPeer as Node, TypeConfig};
use tsoracle_openraft_toolkit::{
BASELINE_WRITE_VERSION, MAX_READABLE_VERSION, MIN_READABLE_VERSION,
};
type NodeId = u64;
pub mod proto {
tonic::include_proto!("tsoracle.raft.peer.v1");
}
use proto::RaftMessage;
use proto::SnapshotChunk;
use proto::SnapshotHeader;
use proto::raft_peer_service_client::RaftPeerServiceClient;
use proto::raft_peer_service_server::{RaftPeerService, RaftPeerServiceServer};
use proto::snapshot_chunk::Kind as ChunkKind;
mod wire {
use super::{BASELINE_WRITE_VERSION, MAX_READABLE_VERSION, MIN_READABLE_VERSION};
pub(super) fn stamp(version: u8) -> u32 {
u32::from(version)
}
pub(super) fn readable_version(format_version: u32) -> Result<u8, String> {
let version = if format_version == 0 {
BASELINE_WRITE_VERSION
} else {
u8::try_from(format_version).map_err(|_| {
format!(
"format_version {format_version} outside readable range \
[{MIN_READABLE_VERSION}, {MAX_READABLE_VERSION}]"
)
})?
};
if !(MIN_READABLE_VERSION..=MAX_READABLE_VERSION).contains(&version) {
return Err(format!(
"format_version {version} outside readable range \
[{MIN_READABLE_VERSION}, {MAX_READABLE_VERSION}]"
));
}
Ok(version)
}
}
pub const SNAPSHOT_CHUNK_SIZE: usize = 1024 * 1024;
pub const MAX_SNAPSHOT_BYTES: usize = 64 * 1024 * 1024;
pub const MAX_PEER_MESSAGE_BYTES: usize = SNAPSHOT_CHUNK_SIZE + 256 * 1024;
const SNAPSHOT_STREAM_TIMEOUT: Duration = Duration::from_secs(60);
type Pool = Arc<Mutex<HashMap<(NodeId, String), RaftPeerServiceClient<Channel>>>>;
pub type WriteVersionSource = Arc<dyn Fn() -> u8 + Send + Sync>;
async fn evict<V>(pool: &Arc<Mutex<HashMap<(NodeId, String), V>>>, target: NodeId, addr: &str) {
pool.lock().await.remove(&(target, addr.to_string()));
}
async fn unary_call<ClientHandle, Body>(
pool: &Arc<Mutex<HashMap<(NodeId, String), ClientHandle>>>,
target: NodeId,
addr: &str,
deadline: Duration,
call: impl Future<Output = Result<tonic::Response<Body>, tonic::Status>>,
) -> Result<Body, RPCError<TypeConfig>> {
match tokio::time::timeout(deadline, call).await {
Ok(Ok(resp)) => Ok(resp.into_inner()),
Ok(Err(status)) => {
evict(pool, target, addr).await;
Err(RPCError::Network(NetworkError::new(&status)))
}
Err(_elapsed) => {
evict(pool, target, addr).await;
let timed_out = std::io::Error::new(
std::io::ErrorKind::TimedOut,
format!("peer RPC exceeded {deadline:?} deadline"),
);
Err(RPCError::Unreachable(Unreachable::new(&timed_out)))
}
}
}
pub struct PeerFactory {
pool: Pool,
tls: Option<ClientTlsConfig>,
active_write_version: WriteVersionSource,
}
impl PeerFactory {
pub fn new(tls: Option<ClientTlsConfig>, active_write_version: WriteVersionSource) -> Self {
Self {
pool: Arc::new(Mutex::new(HashMap::new())),
tls,
active_write_version,
}
}
}
impl RaftNetworkFactory<TypeConfig> for PeerFactory {
type Network = PeerNetwork;
async fn new_client(&mut self, target: NodeId, node: &Node) -> Self::Network {
PeerNetwork {
target,
addr: node.addr.clone(),
pool: self.pool.clone(),
tls: self.tls.clone(),
active_write_version: self.active_write_version.clone(),
}
}
}
pub struct PeerNetwork {
target: NodeId,
addr: String,
pool: Pool,
tls: Option<ClientTlsConfig>,
active_write_version: WriteVersionSource,
}
const CAPABILITY_QUERY_TIMEOUT: Duration = Duration::from_secs(10);
pub struct PeerCapabilitySource {
pool: Pool,
tls: Option<ClientTlsConfig>,
}
impl PeerCapabilitySource {
#[allow(dead_code)]
pub fn new(tls: Option<ClientTlsConfig>) -> Self {
Self {
pool: Arc::new(Mutex::new(HashMap::new())),
tls,
}
}
}
#[async_trait::async_trait]
impl tsoracle_driver_openraft::CapabilitySource for PeerCapabilitySource {
type Node = Node;
async fn query(&self, node_id: NodeId, member: &Node) -> Result<NodeCapabilities, String> {
let network = PeerNetwork {
target: node_id,
addr: member.addr.clone(),
pool: self.pool.clone(),
tls: self.tls.clone(),
active_write_version: Arc::new(|| tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION),
};
network
.capabilities(CAPABILITY_QUERY_TIMEOUT)
.await
.map_err(|err| format!("{err:?}"))
}
}
fn capabilities_response(active_write_version: u8) -> Vec<u8> {
let capabilities = NodeCapabilities::local(active_write_version);
postcard::to_stdvec(&capabilities).unwrap_or_default()
}
impl PeerNetwork {
async fn client(&self) -> Result<RaftPeerServiceClient<Channel>, RPCError<TypeConfig>> {
let key = (self.target, self.addr.clone());
{
let pool = self.pool.lock().await;
if let Some(client) = pool.get(&key) {
return Ok(client.clone());
}
}
let channel = match &self.tls {
Some(tls) => Channel::from_shared(format!("https://{}", self.addr))
.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))?
.tls_config(tls.clone())
.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))?
.connect()
.await
.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))?,
None => Channel::from_shared(format!("http://{}", self.addr))
.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))?
.connect()
.await
.map_err(|err| RPCError::Unreachable(Unreachable::new(&err)))?,
};
let client = RaftPeerServiceClient::new(channel);
self.pool.lock().await.insert(key, client.clone());
Ok(client)
}
pub async fn capabilities(
&self,
deadline: Duration,
) -> Result<NodeCapabilities, RPCError<TypeConfig>> {
let mut client = self.client().await?;
let reply = unary_call(
&self.pool,
self.target,
&self.addr,
deadline,
client.capabilities(RaftMessage {
payload: Vec::new(),
format_version: wire::stamp((self.active_write_version)()),
}),
)
.await?;
let capabilities: NodeCapabilities = postcard::from_bytes(&reply.payload)
.map_err(|err| RPCError::Network(NetworkError::new(&err)))?;
Ok(capabilities)
}
}
impl RaftNetworkV2<TypeConfig> for PeerNetwork {
async fn append_entries(
&mut self,
rpc: AppendEntriesRequest<TypeConfig>,
option: RPCOption,
) -> Result<AppendEntriesResponse<TypeConfig>, RPCError<TypeConfig>> {
let mut c = self.client().await?;
let payload =
postcard::to_stdvec(&rpc).map_err(|err| RPCError::Network(NetworkError::new(&err)))?;
let reply = unary_call(
&self.pool,
self.target,
&self.addr,
option.hard_ttl(),
c.append_entries(RaftMessage {
payload,
format_version: wire::stamp((self.active_write_version)()),
}),
)
.await?;
let _version = wire::readable_version(reply.format_version)
.map_err(|err| RPCError::Network(NetworkError::new(&std::io::Error::other(err))))?;
let body: AppendEntriesResponse<TypeConfig> = postcard::from_bytes(&reply.payload)
.map_err(|err| RPCError::Network(NetworkError::new(&err)))?;
Ok(body)
}
async fn transfer_leader(
&mut self,
req: TransferLeaderRequest<TypeConfig>,
option: RPCOption,
) -> Result<(), RPCError<TypeConfig>> {
let mut c = self.client().await?;
let payload =
postcard::to_stdvec(&req).map_err(|err| RPCError::Network(NetworkError::new(&err)))?;
unary_call(
&self.pool,
self.target,
&self.addr,
option.hard_ttl(),
c.transfer_leader(RaftMessage {
payload,
format_version: wire::stamp((self.active_write_version)()),
}),
)
.await?;
Ok(())
}
async fn vote(
&mut self,
rpc: VoteRequest<TypeConfig>,
option: RPCOption,
) -> Result<VoteResponse<TypeConfig>, RPCError<TypeConfig>> {
let mut c = self.client().await?;
let payload =
postcard::to_stdvec(&rpc).map_err(|err| RPCError::Network(NetworkError::new(&err)))?;
let reply = unary_call(
&self.pool,
self.target,
&self.addr,
option.hard_ttl(),
c.vote(RaftMessage {
payload,
format_version: wire::stamp((self.active_write_version)()),
}),
)
.await?;
let _version = wire::readable_version(reply.format_version)
.map_err(|err| RPCError::Network(NetworkError::new(&std::io::Error::other(err))))?;
let body: VoteResponse<TypeConfig> = postcard::from_bytes(&reply.payload)
.map_err(|err| RPCError::Network(NetworkError::new(&err)))?;
Ok(body)
}
async fn full_snapshot(
&mut self,
vote: VoteOf<TypeConfig>,
snapshot: SnapshotOf<TypeConfig>,
cancel: impl Future<Output = ReplicationClosed> + openraft::OptionalSend + 'static,
_option: RPCOption,
) -> Result<SnapshotResponse<TypeConfig>, StreamingError<TypeConfig>> {
let vote_bytes = postcard::to_stdvec(&vote)
.map_err(|e| StreamingError::Network(NetworkError::new(&e)))?;
let meta_bytes = postcard::to_stdvec(&snapshot.meta)
.map_err(|e| StreamingError::Network(NetworkError::new(&e)))?;
let data_bytes = snapshot.snapshot.into_inner();
let header_chunk = SnapshotChunk {
kind: Some(ChunkKind::Header(SnapshotHeader {
vote: vote_bytes,
meta: meta_bytes,
format_version: wire::stamp((self.active_write_version)()),
})),
};
let data_chunks = data_bytes
.chunks(SNAPSHOT_CHUNK_SIZE)
.map(|c| SnapshotChunk {
kind: Some(ChunkKind::Data(c.to_vec())),
})
.collect::<Vec<_>>();
let outbound =
futures::stream::iter(std::iter::once(header_chunk).chain(data_chunks.into_iter()));
let mut c = self.client().await.map_err(|e| match e {
RPCError::Network(n) => StreamingError::Network(n),
RPCError::Unreachable(u) => StreamingError::Unreachable(u),
RPCError::Timeout(t) => StreamingError::Timeout(t),
})?;
tokio::select! {
result = c.snapshot(outbound) => {
let raw = match result {
Ok(resp) => resp,
Err(err) => {
evict(&self.pool, self.target, &self.addr).await;
return Err(StreamingError::Network(NetworkError::new(&err)));
}
};
let inner = raw.into_inner();
let _version = wire::readable_version(inner.format_version).map_err(|err| {
StreamingError::Network(NetworkError::new(&std::io::Error::other(err)))
})?;
let resp: SnapshotResponse<TypeConfig> = postcard::from_bytes(&inner.payload)
.map_err(|err| StreamingError::Network(NetworkError::new(&err)))?;
Ok(resp)
}
closed = cancel => {
Err(StreamingError::Closed(closed))
}
}
}
}
pub struct PeerServiceImpl<SM = ()> {
pub raft: openraft::Raft<TypeConfig, SM>,
pub active_write_version: WriteVersionSource,
}
#[tonic::async_trait]
impl<SM: Send + Sync + 'static> RaftPeerService for PeerServiceImpl<SM> {
async fn append_entries(
&self,
request: tonic::Request<RaftMessage>,
) -> Result<tonic::Response<RaftMessage>, tonic::Status> {
let message = request.into_inner();
let _version = wire::readable_version(message.format_version)
.map_err(tonic::Status::invalid_argument)?;
let body: AppendEntriesRequest<TypeConfig> = postcard::from_bytes(&message.payload)
.map_err(|e| tonic::Status::invalid_argument(e.to_string()))?;
let resp = self
.raft
.append_entries(body)
.await
.map_err(|e| tonic::Status::internal(e.to_string()))?;
let payload =
postcard::to_stdvec(&resp).map_err(|e| tonic::Status::internal(e.to_string()))?;
Ok(tonic::Response::new(RaftMessage {
payload,
format_version: wire::stamp((self.active_write_version)()),
}))
}
async fn vote(
&self,
request: tonic::Request<RaftMessage>,
) -> Result<tonic::Response<RaftMessage>, tonic::Status> {
let message = request.into_inner();
let _version = wire::readable_version(message.format_version)
.map_err(tonic::Status::invalid_argument)?;
let body: VoteRequest<TypeConfig> = postcard::from_bytes(&message.payload)
.map_err(|e| tonic::Status::invalid_argument(e.to_string()))?;
let resp = self
.raft
.vote(body)
.await
.map_err(|e| tonic::Status::internal(e.to_string()))?;
let payload =
postcard::to_stdvec(&resp).map_err(|e| tonic::Status::internal(e.to_string()))?;
Ok(tonic::Response::new(RaftMessage {
payload,
format_version: wire::stamp((self.active_write_version)()),
}))
}
async fn transfer_leader(
&self,
request: tonic::Request<RaftMessage>,
) -> Result<tonic::Response<RaftMessage>, tonic::Status> {
let message = request.into_inner();
let _version = wire::readable_version(message.format_version)
.map_err(tonic::Status::invalid_argument)?;
let body: TransferLeaderRequest<TypeConfig> = postcard::from_bytes(&message.payload)
.map_err(|e| tonic::Status::invalid_argument(e.to_string()))?;
self.raft
.handle_transfer_leader(body)
.await
.map_err(|e| tonic::Status::internal(e.to_string()))?;
Ok(tonic::Response::new(RaftMessage {
payload: Vec::new(),
format_version: wire::stamp((self.active_write_version)()),
}))
}
async fn snapshot(
&self,
request: tonic::Request<tonic::Streaming<SnapshotChunk>>,
) -> Result<tonic::Response<RaftMessage>, tonic::Status> {
let assembled = tokio::time::timeout(
SNAPSHOT_STREAM_TIMEOUT,
reassemble_snapshot(request.into_inner(), MAX_SNAPSHOT_BYTES),
)
.await
.map_err(|_| tonic::Status::deadline_exceeded("snapshot stream timed out"))??;
let snapshot = openraft::storage::Snapshot {
meta: assembled.meta,
snapshot: Cursor::new(assembled.data),
};
let resp = self
.raft
.install_full_snapshot(assembled.vote, snapshot)
.await
.map_err(|e| tonic::Status::internal(e.to_string()))?;
let payload =
postcard::to_stdvec(&resp).map_err(|e| tonic::Status::internal(e.to_string()))?;
Ok(tonic::Response::new(RaftMessage {
payload,
format_version: wire::stamp((self.active_write_version)()),
}))
}
async fn capabilities(
&self,
_request: tonic::Request<RaftMessage>,
) -> Result<tonic::Response<RaftMessage>, tonic::Status> {
let payload = capabilities_response((self.active_write_version)());
Ok(tonic::Response::new(RaftMessage {
payload,
format_version: wire::stamp((self.active_write_version)()),
}))
}
}
#[derive(Debug)]
struct AssembledSnapshot {
vote: VoteOf<TypeConfig>,
meta: openraft::type_config::alias::SnapshotMetaOf<TypeConfig>,
data: Vec<u8>,
}
async fn reassemble_snapshot<S>(
mut stream: S,
max_bytes: usize,
) -> Result<AssembledSnapshot, tonic::Status>
where
S: futures::Stream<Item = Result<SnapshotChunk, tonic::Status>> + Unpin,
{
let first = stream
.next()
.await
.ok_or_else(|| tonic::Status::invalid_argument("snapshot stream ended before header"))?
.map_err(|e| tonic::Status::internal(format!("snapshot stream error: {e}")))?;
let header = match first.kind {
Some(ChunkKind::Header(h)) => h,
Some(ChunkKind::Data(_)) => {
return Err(tonic::Status::invalid_argument(
"first snapshot chunk must be a header",
));
}
None => {
return Err(tonic::Status::invalid_argument(
"snapshot chunk missing kind",
));
}
};
let _version =
wire::readable_version(header.format_version).map_err(tonic::Status::invalid_argument)?;
let vote: VoteOf<TypeConfig> = postcard::from_bytes(&header.vote)
.map_err(|e| tonic::Status::invalid_argument(format!("bad vote: {e}")))?;
let meta: openraft::type_config::alias::SnapshotMetaOf<TypeConfig> =
postcard::from_bytes(&header.meta)
.map_err(|e| tonic::Status::invalid_argument(format!("bad meta: {e}")))?;
let mut data: Vec<u8> = Vec::new();
while let Some(chunk) = stream.next().await {
let chunk =
chunk.map_err(|e| tonic::Status::internal(format!("snapshot stream error: {e}")))?;
match chunk.kind {
Some(ChunkKind::Data(bytes)) => {
if data.len() + bytes.len() > max_bytes {
return Err(tonic::Status::resource_exhausted(format!(
"snapshot exceeds {max_bytes}-byte reassembly limit"
)));
}
data.extend_from_slice(&bytes);
}
Some(ChunkKind::Header(_)) => {
return Err(tonic::Status::invalid_argument(
"unexpected header chunk after first",
));
}
None => {
return Err(tonic::Status::invalid_argument(
"snapshot chunk missing kind",
));
}
}
}
Ok(AssembledSnapshot { vote, meta, data })
}
pub fn server<SM: Send + Sync + 'static>(
raft: openraft::Raft<TypeConfig, SM>,
active_write_version: WriteVersionSource,
) -> RaftPeerServiceServer<PeerServiceImpl<SM>> {
RaftPeerServiceServer::new(PeerServiceImpl {
raft,
active_write_version,
})
}
#[cfg(test)]
mod tls_tests {
use super::*;
use crate::config::PeerTlsConfig;
use crate::peer_tls::build_peer_tls;
use std::sync::Arc;
use tonic::transport::{Certificate, ClientTlsConfig, Identity};
struct Certs {
ca_pem: String,
node_cert: String,
node_key: String,
other_leaf_cert: String,
other_leaf_key: String,
}
fn mint() -> Certs {
use rcgen::{BasicConstraints, CertificateParams, IsCa, KeyPair};
let mk_ca = |name: &str| {
let key = KeyPair::generate().unwrap();
let mut p = CertificateParams::new(vec![name.to_string()]).unwrap();
p.is_ca = IsCa::Ca(BasicConstraints::Unconstrained);
let cert = p.self_signed(&key).unwrap();
(cert, key)
};
let (ca, ca_key) = mk_ca("tso-ca");
let leaf_key = KeyPair::generate().unwrap();
let leaf_params =
CertificateParams::new(vec!["localhost".to_string(), "127.0.0.1".to_string()]).unwrap();
let leaf = leaf_params.signed_by(&leaf_key, &ca, &ca_key).unwrap();
let (other_ca, other_ca_key) = mk_ca("other-ca");
let other_key = KeyPair::generate().unwrap();
let other_params = CertificateParams::new(vec!["127.0.0.1".to_string()]).unwrap();
let other_leaf = other_params
.signed_by(&other_key, &other_ca, &other_ca_key)
.unwrap();
Certs {
ca_pem: ca.pem(),
node_cert: leaf.pem(),
node_key: leaf_key.serialize_pem(),
other_leaf_cert: other_leaf.pem(),
other_leaf_key: other_key.serialize_pem(),
}
}
fn node_material(c: &Certs, dir: &std::path::Path) -> crate::peer_tls::PeerTlsMaterial {
let cert = dir.join("n.crt");
let key = dir.join("n.key");
let ca = dir.join("ca.crt");
std::fs::write(&cert, &c.node_cert).unwrap();
std::fs::write(&key, &c.node_key).unwrap();
std::fs::write(&ca, &c.ca_pem).unwrap();
build_peer_tls(&PeerTlsConfig { cert, key, ca }).unwrap()
}
#[derive(Clone)]
struct Stub;
#[tonic::async_trait]
impl proto::raft_peer_service_server::RaftPeerService for Stub {
async fn append_entries(
&self,
_: tonic::Request<proto::RaftMessage>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Err(tonic::Status::unimplemented("stub"))
}
async fn vote(
&self,
_: tonic::Request<proto::RaftMessage>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Err(tonic::Status::unimplemented("stub"))
}
async fn snapshot(
&self,
_: tonic::Request<tonic::Streaming<proto::SnapshotChunk>>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Err(tonic::Status::unimplemented("stub"))
}
async fn transfer_leader(
&self,
_: tonic::Request<proto::RaftMessage>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Err(tonic::Status::unimplemented("stub"))
}
async fn capabilities(
&self,
_: tonic::Request<proto::RaftMessage>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Err(tonic::Status::unimplemented("stub"))
}
}
async fn spawn_stub(server_tls: tonic::transport::ServerTlsConfig) -> std::net::SocketAddr {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
tonic::transport::Server::builder()
.tls_config(server_tls)
.unwrap()
.add_service(proto::raft_peer_service_server::RaftPeerServiceServer::new(
Stub,
))
.serve_with_incoming(tokio_stream::wrappers::TcpListenerStream::new(listener))
.await
.ok();
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
addr
}
fn make_net(addr: std::net::SocketAddr, tls: Option<ClientTlsConfig>) -> PeerNetwork {
PeerNetwork {
target: 2,
addr: addr.to_string(),
pool: Arc::new(Mutex::new(HashMap::new())),
tls,
active_write_version: Arc::new(|| tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION),
}
}
async fn probe(net: PeerNetwork) -> tonic::Code {
match net.client().await {
Err(_) => tonic::Code::Unavailable,
Ok(mut c) => {
match c
.append_entries(tonic::Request::new(proto::RaftMessage {
payload: Vec::new(),
format_version: 0,
}))
.await
{
Ok(_) => tonic::Code::Ok,
Err(s) => s.code(),
}
}
}
}
#[tokio::test]
async fn valid_node_cert_connects() {
let dir = tempfile::tempdir().unwrap();
let c = mint();
let m = node_material(&c, dir.path());
let addr = spawn_stub(m.server.clone()).await;
assert_eq!(
probe(make_net(addr, Some(m.client.clone()))).await,
tonic::Code::Unimplemented
);
}
#[tokio::test]
async fn no_client_cert_rejected() {
let dir = tempfile::tempdir().unwrap();
let c = mint();
let m = node_material(&c, dir.path());
let addr = spawn_stub(m.server.clone()).await;
let no_id = ClientTlsConfig::new()
.ca_certificate(Certificate::from_pem(&c.ca_pem))
.domain_name("localhost");
let code = probe(make_net(addr, Some(no_id))).await;
assert_ne!(
code,
tonic::Code::Unimplemented,
"server must reject a client with no cert"
);
}
#[tokio::test]
async fn wrong_ca_client_rejected() {
let dir = tempfile::tempdir().unwrap();
let c = mint();
let m = node_material(&c, dir.path());
let addr = spawn_stub(m.server.clone()).await;
let wrong = ClientTlsConfig::new()
.ca_certificate(Certificate::from_pem(&c.ca_pem))
.identity(Identity::from_pem(&c.other_leaf_cert, &c.other_leaf_key))
.domain_name("localhost");
let code = probe(make_net(addr, Some(wrong))).await;
assert_ne!(
code,
tonic::Code::Unimplemented,
"server must reject a cert from a foreign CA"
);
}
#[tokio::test]
async fn plaintext_against_tls_fails() {
let dir = tempfile::tempdir().unwrap();
let c = mint();
let m = node_material(&c, dir.path());
let addr = spawn_stub(m.server.clone()).await;
let code = probe(make_net(addr, None)).await;
assert_ne!(
code,
tonic::Code::Unimplemented,
"plaintext must not reach a TLS-only server"
);
}
}
#[cfg(test)]
mod tests {
use std::collections::HashMap;
use std::sync::Arc;
use openraft::Vote;
use openraft::type_config::alias::SnapshotMetaOf;
use tokio::sync::Mutex;
use super::*;
#[tokio::test]
async fn pool_key_distinguishes_addr_changes() {
let mut map: HashMap<(u64, String), u8> = HashMap::new();
map.insert((1, "old:1".to_string()), 0);
assert!(!map.contains_key(&(1, "new:1".to_string())));
assert!(map.contains_key(&(1, "old:1".to_string())));
}
#[tokio::test]
async fn evict_removes_only_the_targeted_entry() {
let pool: Arc<Mutex<HashMap<(u64, String), u8>>> = Arc::new(Mutex::new(HashMap::new()));
{
let mut guard = pool.lock().await;
guard.insert((1, "a:1".to_string()), 0);
guard.insert((2, "b:2".to_string()), 0);
}
evict(&pool, 1, "a:1").await;
let guard = pool.lock().await;
assert!(!guard.contains_key(&(1, "a:1".to_string())));
assert!(guard.contains_key(&(2, "b:2".to_string())));
}
fn seeded_pool() -> Arc<Mutex<HashMap<(u64, String), u8>>> {
let pool = Arc::new(Mutex::new(HashMap::new()));
pool.try_lock().unwrap().insert((1, "a:1".to_string()), 0);
pool
}
#[tokio::test]
async fn unary_call_deadline_elapse_evicts_and_reports_unreachable() {
let pool = seeded_pool();
let never = std::future::pending::<Result<tonic::Response<u8>, tonic::Status>>();
let err = unary_call(&pool, 1, "a:1", Duration::from_millis(10), never)
.await
.expect_err("a never-resolving call must hit the deadline");
assert!(matches!(err, RPCError::Unreachable(_)));
assert!(!pool.lock().await.contains_key(&(1, "a:1".to_string())));
}
#[tokio::test]
async fn unary_call_transport_error_evicts_and_reports_network() {
let pool = seeded_pool();
let failed = async { Err::<tonic::Response<u8>, _>(tonic::Status::unavailable("down")) };
let err = unary_call(&pool, 1, "a:1", Duration::from_secs(5), failed)
.await
.expect_err("a transport error must propagate");
assert!(matches!(err, RPCError::Network(_)));
assert!(!pool.lock().await.contains_key(&(1, "a:1".to_string())));
}
#[tokio::test]
async fn unary_call_success_returns_body_and_keeps_client() {
let pool = seeded_pool();
let ok = async { Ok(tonic::Response::new(42u8)) };
let body = unary_call(&pool, 1, "a:1", Duration::from_secs(5), ok)
.await
.expect("a successful call returns its body");
assert_eq!(body, 42);
assert!(pool.lock().await.contains_key(&(1, "a:1".to_string())));
}
fn header_chunk() -> SnapshotChunk {
let vote: VoteOf<TypeConfig> = Vote::new(1, 1);
let meta = SnapshotMetaOf::<TypeConfig> {
last_log_id: None,
last_membership: Default::default(),
snapshot_id: "test-snap".to_string(),
};
SnapshotChunk {
kind: Some(ChunkKind::Header(SnapshotHeader {
vote: postcard::to_stdvec(&vote).expect("encode vote"),
meta: postcard::to_stdvec(&meta).expect("encode meta"),
format_version: wire::stamp(tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION),
})),
}
}
fn data_chunk(bytes: &[u8]) -> SnapshotChunk {
SnapshotChunk {
kind: Some(ChunkKind::Data(bytes.to_vec())),
}
}
fn ok_stream(
chunks: Vec<SnapshotChunk>,
) -> impl futures::Stream<Item = Result<SnapshotChunk, tonic::Status>> + Unpin {
futures::stream::iter(chunks.into_iter().map(Ok))
}
#[tokio::test]
async fn snapshot_over_limit_is_resource_exhausted() {
let chunks = vec![
header_chunk(),
data_chunk(&[0u8; 600]),
data_chunk(&[0u8; 600]),
];
let err = reassemble_snapshot(ok_stream(chunks), 1024)
.await
.expect_err("over-limit stream must be rejected");
assert_eq!(err.code(), tonic::Code::ResourceExhausted);
}
#[tokio::test]
async fn snapshot_under_limit_assembles() {
let chunks = vec![header_chunk(), data_chunk(b"hello "), data_chunk(b"world")];
let assembled = reassemble_snapshot(ok_stream(chunks), 1024)
.await
.expect("under-limit stream assembles");
assert_eq!(assembled.data, b"hello world");
}
#[tokio::test]
async fn data_before_header_is_invalid_argument() {
let chunks = vec![data_chunk(b"premature")];
let err = reassemble_snapshot(ok_stream(chunks), 1024)
.await
.expect_err("data before header must be rejected");
assert_eq!(err.code(), tonic::Code::InvalidArgument);
}
#[test]
fn absent_format_version_reads_as_baseline() {
let version = wire::readable_version(0).expect("0 normalizes to baseline");
assert_eq!(version, tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION);
}
#[test]
fn in_range_format_version_passes_through() {
let stamped = tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION;
let version = wire::readable_version(u32::from(stamped)).expect("in-range");
assert_eq!(version, stamped);
}
#[test]
fn out_of_range_format_version_is_rejected() {
let too_new = u32::from(tsoracle_openraft_toolkit::MAX_READABLE_VERSION) + 1;
let err = wire::readable_version(too_new).expect_err("out-of-range rejected");
assert!(
err.contains("format_version"),
"message names the field: {err}"
);
}
#[test]
fn stamp_widens_to_u32() {
assert_eq!(wire::stamp(3), 3u32);
assert_eq!(wire::stamp(255), 255u32);
}
#[tokio::test]
async fn peer_network_holds_the_active_write_version_source() {
let factory_version = Arc::new(std::sync::atomic::AtomicU8::new(
tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION,
));
let version_for_source = factory_version.clone();
let mut factory = PeerFactory::new(
None,
Arc::new(move || version_for_source.load(std::sync::atomic::Ordering::Relaxed)),
);
let node = Node {
addr: "127.0.0.1:1".to_string(),
service_endpoint: String::new(),
admin_endpoint: String::new(),
};
let net = factory.new_client(7, &node).await;
assert_eq!(
(net.active_write_version)(),
tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION
);
factory_version.store(5, std::sync::atomic::Ordering::Relaxed);
assert_eq!((net.active_write_version)(), 5);
}
mod test_support {
use super::*;
use openraft::{Config, Raft};
use rocksdb::{ColumnFamilyDescriptor, DB, Options};
use std::collections::BTreeMap;
use tempfile::TempDir;
use tsoracle_driver_openraft::{
HighWaterStateMachine, OpenraftLogCodec, OpenraftPeer, RocksdbSnapshotStore,
SnapshotStore,
};
use tsoracle_openraft_toolkit::{ActiveWriteVersion, Flat, RocksdbLogStore};
pub(super) async fn single_node_raft() -> (Raft<TypeConfig, HighWaterStateMachine>, TempDir)
{
let dir = tempfile::tempdir().expect("tempdir");
let mut opts = Options::default();
opts.create_if_missing(true);
opts.create_missing_column_families(true);
let cfs = vec![
ColumnFamilyDescriptor::new("raft_log", Options::default()),
ColumnFamilyDescriptor::new("raft_meta", Options::default()),
ColumnFamilyDescriptor::new("raft_snapshot", Options::default()),
];
let db =
Arc::new(DB::open_cf_descriptors(&opts, dir.path(), cfs).expect("open rocksdb"));
let cell = ActiveWriteVersion::default();
let log_store: RocksdbLogStore<TypeConfig, Flat, OpenraftLogCodec> =
RocksdbLogStore::open(db.clone(), "raft_log", "raft_meta", Flat)
.expect("open log store")
.with_active_write_version(cell.clone());
let snapshot_store: Arc<dyn SnapshotStore> = Arc::new(
RocksdbSnapshotStore::open(db, "raft_snapshot").expect("open snapshot store"),
);
let state_machine =
HighWaterStateMachine::with_store_and_active_version(snapshot_store, cell)
.expect("state machine");
let config = Arc::new(
Config {
heartbeat_interval: 50,
election_timeout_min: 150,
election_timeout_max: 300,
..Default::default()
}
.validate()
.expect("validate config"),
);
let version_source: WriteVersionSource =
Arc::new(|| tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION);
let network = PeerFactory::new(None, version_source);
let raft = Raft::<TypeConfig, HighWaterStateMachine>::new(
1,
config,
network,
log_store,
state_machine,
)
.await
.expect("raft new");
let mut members: BTreeMap<u64, OpenraftPeer> = BTreeMap::new();
members.insert(
1,
OpenraftPeer {
addr: "127.0.0.1:1".to_string(),
service_endpoint: String::new(),
admin_endpoint: String::new(),
},
);
let _ = raft.initialize(members).await;
(raft, dir)
}
}
#[tokio::test]
async fn vote_round_trips_with_baseline_format_version() {
use openraft::Vote;
let (raft, _temp) = test_support::single_node_raft().await;
let version_source: WriteVersionSource =
Arc::new(|| tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let service = server(raft, version_source.clone());
tokio::spawn(async move {
tonic::transport::Server::builder()
.add_service(service)
.serve_with_incoming(tokio_stream::wrappers::TcpListenerStream::new(listener))
.await
.ok();
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let mut net = PeerNetwork {
target: 1,
addr: addr.to_string(),
pool: Arc::new(Mutex::new(HashMap::new())),
tls: None,
active_write_version: version_source,
};
let request = openraft::raft::VoteRequest::<TypeConfig>::new(Vote::new(1, 1), None);
let response = net
.vote(request, RPCOption::new(Duration::from_secs(5)))
.await
.expect("vote round-trips at baseline");
let _ = response;
}
#[tokio::test]
async fn server_rejects_out_of_range_format_version() {
let (raft, _temp) = test_support::single_node_raft().await;
let version_source: WriteVersionSource =
Arc::new(|| tsoracle_openraft_toolkit::BASELINE_WRITE_VERSION);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let service = server(raft, version_source);
tokio::spawn(async move {
tonic::transport::Server::builder()
.add_service(service)
.serve_with_incoming(tokio_stream::wrappers::TcpListenerStream::new(listener))
.await
.ok();
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let mut client = proto::raft_peer_service_client::RaftPeerServiceClient::connect(format!(
"http://{addr}"
))
.await
.unwrap();
let too_new = u32::from(tsoracle_openraft_toolkit::MAX_READABLE_VERSION) + 1;
let status = client
.vote(tonic::Request::new(proto::RaftMessage {
payload: Vec::new(),
format_version: too_new,
}))
.await
.expect_err("an out-of-range format_version must be refused");
assert_eq!(status.code(), tonic::Code::InvalidArgument);
}
#[tokio::test]
async fn snapshot_header_format_version_is_read_and_range_checked() {
let ok_chunks = vec![header_chunk(), data_chunk(b"snap")];
let assembled = reassemble_snapshot(ok_stream(ok_chunks), 1024)
.await
.expect("baseline-framed header assembles");
assert_eq!(assembled.data, b"snap");
let vote: VoteOf<TypeConfig> = openraft::Vote::new(1, 1);
let meta = SnapshotMetaOf::<TypeConfig> {
last_log_id: None,
last_membership: Default::default(),
snapshot_id: "bad".to_string(),
};
let bad_header = SnapshotChunk {
kind: Some(ChunkKind::Header(SnapshotHeader {
vote: postcard::to_stdvec(&vote).unwrap(),
meta: postcard::to_stdvec(&meta).unwrap(),
format_version: u32::from(tsoracle_openraft_toolkit::MAX_READABLE_VERSION) + 1,
})),
};
let err = reassemble_snapshot(ok_stream(vec![bad_header]), 1024)
.await
.expect_err("out-of-range header rejected");
assert_eq!(err.code(), tonic::Code::InvalidArgument);
}
#[tokio::test]
async fn snapshot_header_absent_format_version_reads_as_baseline() {
let vote: VoteOf<TypeConfig> = Vote::new(1, 1);
let meta = SnapshotMetaOf::<TypeConfig> {
last_log_id: None,
last_membership: Default::default(),
snapshot_id: "legacy".to_string(),
};
let legacy_header = SnapshotChunk {
kind: Some(ChunkKind::Header(SnapshotHeader {
vote: postcard::to_stdvec(&vote).unwrap(),
meta: postcard::to_stdvec(&meta).unwrap(),
format_version: 0,
})),
};
let assembled = reassemble_snapshot(ok_stream(vec![legacy_header, data_chunk(b"x")]), 1024)
.await
.expect("absent format_version assembles as baseline");
assert_eq!(assembled.data, b"x");
}
#[test]
fn capabilities_response_reports_local_node() {
let payload = capabilities_response(7);
let decoded: NodeCapabilities =
postcard::from_bytes(&payload).expect("decode NodeCapabilities");
assert_eq!(
decoded,
NodeCapabilities {
min_readable_version: tsoracle_openraft_toolkit::MIN_READABLE_VERSION,
max_readable_version: tsoracle_openraft_toolkit::MAX_READABLE_VERSION,
active_write_version: 7,
}
);
}
#[test]
fn capabilities_payload_round_trips_client_side() {
let server_payload = capabilities_response(4);
let message = RaftMessage {
payload: server_payload,
format_version: 0,
};
let decoded: NodeCapabilities =
postcard::from_bytes(&message.payload).expect("client decode");
assert_eq!(decoded.active_write_version, 4);
}
#[tokio::test]
async fn peer_capability_source_round_trips_against_live_server() {
use proto::raft_peer_service_server::{
RaftPeerService as ProtoService, RaftPeerServiceServer,
};
use tsoracle_driver_openraft::CapabilitySource;
struct CapStub {
active_write_version: u8,
}
#[tonic::async_trait]
impl ProtoService for CapStub {
async fn append_entries(
&self,
_: tonic::Request<proto::RaftMessage>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Err(tonic::Status::unimplemented("cap stub"))
}
async fn vote(
&self,
_: tonic::Request<proto::RaftMessage>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Err(tonic::Status::unimplemented("cap stub"))
}
async fn snapshot(
&self,
_: tonic::Request<tonic::Streaming<proto::SnapshotChunk>>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Err(tonic::Status::unimplemented("cap stub"))
}
async fn transfer_leader(
&self,
_: tonic::Request<proto::RaftMessage>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Err(tonic::Status::unimplemented("cap stub"))
}
async fn capabilities(
&self,
_: tonic::Request<proto::RaftMessage>,
) -> Result<tonic::Response<proto::RaftMessage>, tonic::Status> {
Ok(tonic::Response::new(proto::RaftMessage {
payload: capabilities_response(self.active_write_version),
format_version: 0,
}))
}
}
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
tonic::transport::Server::builder()
.add_service(RaftPeerServiceServer::new(CapStub {
active_write_version: 5,
}))
.serve_with_incoming(tokio_stream::wrappers::TcpListenerStream::new(listener))
.await
.ok();
});
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
let source = PeerCapabilitySource::new(None);
let node = Node {
addr: addr.to_string(),
service_endpoint: String::new(),
admin_endpoint: String::new(),
};
let capabilities = source
.query(2, &node)
.await
.expect("live capabilities query");
assert_eq!(capabilities.active_write_version, 5);
assert_eq!(
capabilities.min_readable_version,
tsoracle_openraft_toolkit::MIN_READABLE_VERSION
);
assert_eq!(
capabilities.max_readable_version,
tsoracle_openraft_toolkit::MAX_READABLE_VERSION
);
}
}