#![cfg_attr(not(test), warn(clippy::unwrap_used, clippy::expect_used))]
mod cli;
#[cfg(feature = "openraft")]
use std::collections::BTreeMap;
use anyhow::{Context, Result};
use clap::Parser;
use cli::{Cli, Cmd, CommonServeArgs, ServeCmd};
use tracing_subscriber::EnvFilter;
use tsoracle_server::Server;
use tsoracle_standalone::{DriverConfig, Standalone};
#[cfg(any(
not(feature = "file"),
not(feature = "openraft"),
not(feature = "paxos")
))]
fn available_drivers() -> &'static [&'static str] {
&[
#[cfg(feature = "file")]
"file",
#[cfg(feature = "openraft")]
"openraft",
#[cfg(feature = "paxos")]
"paxos",
]
}
#[cfg(any(
not(feature = "file"),
not(feature = "openraft"),
not(feature = "paxos")
))]
fn not_compiled_in(driver: &str) -> anyhow::Error {
anyhow::anyhow!(
"this build does not include the {driver} driver; rebuild with `--features {driver}`. \
available drivers: {}",
available_drivers().join(", ")
)
}
#[tokio::main]
async fn main() -> Result<()> {
let cli = Cli::parse();
match cli.cmd {
Some(Cmd::Init(args)) => run_init(args.state_dir, args.seed_physical_ms),
Some(Cmd::Serve(serve)) => dispatch_serve(*serve).await,
#[cfg(feature = "openraft")]
Some(Cmd::Admin(cmd)) => dispatch_admin(cmd).await,
None => {
#[cfg(feature = "file")]
{
let file = cli.serve_file;
let cfg = DriverConfig::File(tsoracle_standalone::FileConfig {
state_dir: file.state_dir,
});
return run_serve(file.common, cfg).await;
}
#[cfg(not(feature = "file"))]
{
anyhow::bail!(
"no subcommand given and this build excludes the file driver; \
specify `serve <driver>`. available drivers: {}",
available_drivers().join(", ")
);
}
}
}
}
async fn dispatch_serve(serve: ServeCmd) -> Result<()> {
match serve {
ServeCmd::File(args) => {
#[cfg(feature = "file")]
{
let cfg = DriverConfig::File(tsoracle_standalone::FileConfig {
state_dir: args.state_dir,
});
run_serve(args.common, cfg).await
}
#[cfg(not(feature = "file"))]
{
let _ = args;
Err(not_compiled_in("file"))
}
}
ServeCmd::Openraft(args) => {
#[cfg(feature = "openraft")]
{
let members = match args.members {
Some(s) => Some(parse_members(&s)?),
None => None,
};
let cfg = DriverConfig::Openraft(tsoracle_standalone::OpenraftConfig {
id: args.id,
raft_addr: args.raft_addr,
raft_dir: args.raft_dir,
bootstrap: args.bootstrap,
initial_membership: members,
tuning: tsoracle_standalone::RaftTuning {
heartbeat_ms: args.heartbeat_ms,
election_min_ms: args.election_min_ms,
election_max_ms: args.election_max_ms,
},
peer_tls: peer_tls_config(
args.peer_tls_cert,
args.peer_tls_key,
args.peer_tls_ca,
)?,
admin_listen: args.admin_listen,
});
run_serve(args.common, cfg).await
}
#[cfg(not(feature = "openraft"))]
{
let _ = args;
Err(not_compiled_in("openraft"))
}
}
ServeCmd::Paxos(args) => {
#[cfg(feature = "paxos")]
{
let cfg = DriverConfig::Paxos(tsoracle_standalone::PaxosConfig {
node_id: args.node_id,
peer_listen: args.peer_listen,
peers: tsoracle_standalone::parse_peer_map(&args.peers)
.map_err(anyhow::Error::msg)?,
tso_peers: tsoracle_standalone::parse_peer_map(&args.tso_peers)
.map_err(anyhow::Error::msg)?,
data_dir: args.data_dir,
tick_interval: args.tick_interval,
peer_tls: peer_tls_config(
args.peer_tls_cert,
args.peer_tls_key,
args.peer_tls_ca,
)?,
});
run_serve(args.common, cfg).await
}
#[cfg(not(feature = "paxos"))]
{
let _ = args;
Err(not_compiled_in("paxos"))
}
}
}
}
#[cfg(feature = "file")]
fn run_init(state_dir: std::path::PathBuf, seed_physical_ms: u64) -> Result<()> {
tsoracle_standalone::init_file_seeded(&state_dir, seed_physical_ms)
.with_context(|| format!("init state_dir={}", state_dir.display()))?;
println!(
"Initialized {} at seed physical_ms={seed_physical_ms}",
state_dir.display()
);
Ok(())
}
#[cfg(not(feature = "file"))]
fn run_init(_state_dir: std::path::PathBuf, _seed_physical_ms: u64) -> Result<()> {
Err(not_compiled_in("file"))
}
fn client_tls_config(
common: &CommonServeArgs,
) -> anyhow::Result<Option<tonic::transport::ServerTlsConfig>> {
match (&common.tls_cert, &common.tls_key) {
(None, None) => {
if common.tls_client_ca.is_some() {
anyhow::bail!("--tls-client-ca requires --tls-cert and --tls-key");
}
Ok(None)
}
(Some(cert), Some(key)) => {
let cert_pem =
std::fs::read(cert).with_context(|| format!("read {}", cert.display()))?;
let key_pem = std::fs::read(key).with_context(|| format!("read {}", key.display()))?;
let mut tls = tonic::transport::ServerTlsConfig::new()
.identity(tonic::transport::Identity::from_pem(&cert_pem, &key_pem));
if let Some(ca) = &common.tls_client_ca {
let ca_pem = std::fs::read(ca).with_context(|| format!("read {}", ca.display()))?;
tls = tls.client_ca_root(tonic::transport::Certificate::from_pem(&ca_pem));
}
tonic::transport::Server::builder()
.tls_config(tls.clone())
.context("invalid client-API TLS configuration")?;
Ok(Some(tls))
}
_ => anyhow::bail!("--tls-cert and --tls-key must be provided together"),
}
}
#[cfg(any(feature = "openraft", feature = "paxos"))]
fn peer_tls_config(
cert: Option<std::path::PathBuf>,
key: Option<std::path::PathBuf>,
ca: Option<std::path::PathBuf>,
) -> anyhow::Result<Option<tsoracle_standalone::PeerTlsConfig>> {
match (cert, key, ca) {
(None, None, None) => Ok(None),
(Some(cert), Some(key), Some(ca)) => {
Ok(Some(tsoracle_standalone::PeerTlsConfig { cert, key, ca }))
}
_ => anyhow::bail!(
"--peer-tls-cert, --peer-tls-key, and --peer-tls-ca must all be set together"
),
}
}
async fn run_serve(common: CommonServeArgs, cfg: DriverConfig) -> Result<()> {
tracing_subscriber::fmt()
.with_env_filter(EnvFilter::try_new(&common.log).unwrap_or_else(|_| EnvFilter::new("info")))
.init();
let mut node: Standalone = tsoracle_standalone::build(cfg)
.await
.context("driver bootstrap")?;
let drain = node.take_drain();
let tls = client_tls_config(&common)?;
let mut builder = Server::builder()
.consensus_driver(node.driver.clone())
.window_ahead(common.window_ahead)
.failover_advance(common.failover_advance);
if let Some(tls) = tls {
builder = builder.tls_config(tls);
}
let server = builder.build().context("server build")?;
let listener = tokio::net::TcpListener::bind(common.listen)
.await
.with_context(|| format!("bind {}", common.listen))?;
let local_addr = listener.local_addr().context("listener.local_addr()")?;
println!("serving on {local_addr}");
tracing::info!(addr = %local_addr, "tsoracle serving");
let shutdown = async move {
tsoracle_server::shutdown_signal().await;
if let Some(drain) = drain {
drain.await;
}
};
let result = server
.serve_with_listener(listener, shutdown)
.await
.context("serve");
node.shutdown().await;
result
}
#[cfg(feature = "openraft")]
fn parse_members(input: &str) -> Result<BTreeMap<u64, tsoracle_standalone::MemberAddr>> {
let mut out = BTreeMap::new();
for entry in input.split(',') {
let entry = entry.trim();
if entry.is_empty() {
continue;
}
let (id, addrs) = entry.split_once('=').with_context(|| {
format!("bad member {entry:?}, expected id=raft_addr/service_endpoint/admin_endpoint")
})?;
let mut parts = addrs.split('/');
let raft_addr = parts.next().filter(|s| !s.is_empty());
let service_endpoint = parts.next();
let admin_endpoint = parts.next();
let (Some(raft_addr), Some(service_endpoint), Some(admin_endpoint)) =
(raft_addr, service_endpoint, admin_endpoint)
else {
anyhow::bail!(
"bad member {entry:?}, expected raft_addr/service_endpoint/admin_endpoint"
);
};
out.insert(
id.trim()
.parse()
.with_context(|| format!("bad member id in {entry:?}"))?,
tsoracle_standalone::MemberAddr {
raft_addr: raft_addr.trim().to_string(),
service_endpoint: service_endpoint.trim().to_string(),
admin_endpoint: admin_endpoint.trim().to_string(),
},
);
}
Ok(out)
}
#[cfg(feature = "openraft")]
async fn dispatch_admin(cmd: cli::AdminCmd) -> Result<()> {
use cli::AdminCmd;
use tsoracle_standalone::admin_proto::membership_admin_client::MembershipAdminClient;
use tsoracle_standalone::admin_proto::{
AddLearnerRequest, AdminErrorKind, ChangeResponse, ListMembersRequest, MemberRole,
PromoteRequest, RemoveNodeRequest,
};
async fn with_redirect<MakeCall>(endpoint: String, op: MakeCall) -> Result<ChangeResponse>
where
MakeCall: Fn(
MembershipAdminClient<tonic::transport::Channel>,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<ChangeResponse, tonic::Status>> + Send>,
>,
{
let client = MembershipAdminClient::connect(endpoint.clone())
.await
.with_context(|| format!("connect {endpoint}"))?;
let resp = op(client).await.context("admin rpc")?;
if resp.error == AdminErrorKind::NotLeader as i32 && !resp.leader_admin_endpoint.is_empty()
{
let leader = format!("http://{}", resp.leader_admin_endpoint);
let client = MembershipAdminClient::connect(leader.clone())
.await
.with_context(|| format!("connect leader {leader}"))?;
return op(client).await.context("admin rpc (leader)");
}
Ok(resp)
}
fn report(resp: ChangeResponse) -> Result<()> {
if resp.ok {
println!("ok");
Ok(())
} else {
let kind = AdminErrorKind::try_from(resp.error)
.map(|kind| format!("{kind:?}"))
.unwrap_or_else(|_| resp.error.to_string());
anyhow::bail!("admin error ({kind}): {}", resp.message)
}
}
match cmd {
AdminCmd::Members(args) => {
let mut client = MembershipAdminClient::connect(args.endpoint.clone())
.await
.with_context(|| format!("connect {}", args.endpoint))?;
let view = client
.list_members(ListMembersRequest {})
.await
.context("list_members")?
.into_inner();
let leader_str: String = if view.has_leader {
view.leader.to_string()
} else {
"none".into()
};
println!("leader: {leader_str}");
for member in view.members {
let role = MemberRole::try_from(member.role)
.map(|role| format!("{role:?}"))
.unwrap_or_else(|_| member.role.to_string());
println!(
" id={} role={role} raft={} service={} admin={}",
member.id, member.raft_addr, member.service_endpoint, member.admin_endpoint
);
}
Ok(())
}
AdminCmd::AddLearner(args) => report(
with_redirect(args.endpoint.clone(), move |mut client| {
let request = AddLearnerRequest {
id: args.id,
raft_addr: args.raft_addr.clone(),
service_endpoint: args.service_endpoint.clone(),
admin_endpoint: args.admin_endpoint.clone(),
};
Box::pin(async move { client.add_learner(request).await.map(|r| r.into_inner()) })
})
.await?,
),
AdminCmd::Promote(args) => report(
with_redirect(args.endpoint.clone(), move |mut client| {
Box::pin(async move {
client
.promote(PromoteRequest { id: args.id })
.await
.map(|r| r.into_inner())
})
})
.await?,
),
AdminCmd::Remove(args) => report(
with_redirect(args.endpoint.clone(), move |mut client| {
Box::pin(async move {
client
.remove_node(RemoveNodeRequest { id: args.id })
.await
.map(|r| r.into_inner())
})
})
.await?,
),
}
}
#[cfg(all(test, feature = "openraft"))]
mod parse_members_tests {
use super::parse_members;
#[test]
fn parses_three_address_members() {
let map = parse_members("1=10.0.0.1:9/10.0.0.1:8/10.0.0.1:7").unwrap();
let member = map.get(&1).expect("id 1 present");
assert_eq!(member.raft_addr, "10.0.0.1:9");
assert_eq!(member.service_endpoint, "10.0.0.1:8");
assert_eq!(member.admin_endpoint, "10.0.0.1:7");
}
#[test]
fn rejects_member_missing_admin_endpoint() {
let err = parse_members("1=10.0.0.1:9/10.0.0.1:8").unwrap_err();
assert!(
err.to_string()
.contains("raft_addr/service_endpoint/admin_endpoint"),
"got: {err}"
);
}
}