use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use std::time::Duration;
use object_store::ObjectStoreExt;
use object_store::PutPayload;
use tokio::sync::Semaphore;
use tokio::task::JoinHandle;
use tokio::time::Instant;
use tokio_util::sync::CancellationToken;
use tracing::{info, warn};
use crate::config::{ResolvedServices, Service};
use crate::heartbeat::{self, Heartbeat, ServiceSummary};
use crate::root::{ConfigPoller, FleetRoot, FleetState, HeartbeatEntry, node_view};
use crate::services::{self, DatabaseHandle};
use crate::{SLATEDB_VERSION, mirror, placement};
#[derive(Clone, Debug)]
pub struct NodeOptions {
pub node_id: String,
pub services: Vec<Service>,
pub max_mirror_jobs: usize,
pub rclone: Option<String>,
}
impl Default for NodeOptions {
fn default() -> Self {
Self {
node_id: String::new(),
services: Service::ALL.to_vec(),
max_mirror_jobs: std::thread::available_parallelism().map_or(4, |p| p.get()),
rclone: None,
}
}
}
impl NodeOptions {
pub fn new(node_id: impl Into<String>) -> Self {
Self {
node_id: node_id.into(),
..Self::default()
}
}
pub fn with_services(mut self, services: impl IntoIterator<Item = Service>) -> Self {
self.services = services.into_iter().collect();
self
}
pub fn with_max_mirror_jobs(mut self, max_mirror_jobs: usize) -> Self {
self.max_mirror_jobs = max_mirror_jobs.max(1);
self
}
pub fn with_rclone(mut self, rclone: impl Into<String>) -> Self {
self.rclone = Some(rclone.into());
self
}
}
#[derive(Debug, thiserror::Error)]
pub enum DaemonError {
#[error("invalid node id: {0}")]
NodeId(String),
#[error(transparent)]
Root(#[from] crate::root::OpenError),
}
pub type Assignment = (String, Service, Option<String>);
pub fn owned_assignments(
node_id: &str,
services: &[Service],
entries: &[HeartbeatEntry],
state: &FleetState,
) -> HashMap<Assignment, Arc<ResolvedServices>> {
let mut nodes = node_view(entries, state.config.node.heartbeat_timeout.0);
if !nodes.iter().any(|n| n.node_id == node_id) {
nodes.push(crate::root::NodeView {
node_id: node_id.to_string(),
services: services.to_vec(),
age: Duration::ZERO,
});
}
let mut owned = HashMap::new();
for (url, db) in &state.databases {
let resolved = Arc::new(state.config.resolve(Some(db)));
for &service in &resolved.services {
let candidates: Vec<&str> = nodes
.iter()
.filter(|n| n.services.contains(&service))
.map(|n| n.node_id.as_str())
.collect();
if service == Service::Mirror {
for applied in mirror::applied_targets(url, &resolved.mirror) {
if placement::owner_target(url, &applied.name, &candidates) == Some(node_id) {
owned.insert((url.clone(), service, Some(applied.name)), resolved.clone());
}
}
continue;
}
let count = match service {
Service::CompactionWorkers => resolved.workers.count as usize,
_ => 1,
};
let owners = placement::owners(url, service, count, &candidates);
if owners.contains(&node_id) {
owned.insert((url.clone(), service, None), resolved.clone());
}
}
}
owned
}
#[derive(Clone, Copy, PartialEq)]
enum TaskState {
Running,
Backoff,
}
type TaskStates = Arc<Mutex<HashMap<Assignment, (u64, TaskState)>>>;
static NEXT_INSTANCE: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
struct RunningTask {
token: CancellationToken,
handle: JoinHandle<()>,
fingerprint: u64,
}
pub async fn run(
root: FleetRoot,
options: NodeOptions,
shutdown: CancellationToken,
) -> Result<(), DaemonError> {
heartbeat::validate_node_id(&options.node_id).map_err(DaemonError::NodeId)?;
let node = Node {
object_name: heartbeat::object_name(&options.node_id, &options.services),
mirror_jobs: Arc::new(Semaphore::new(options.max_mirror_jobs.max(1))),
states: TaskStates::default(),
options,
root,
};
node.run(shutdown).await;
Ok(())
}
struct Node {
options: NodeOptions,
object_name: String,
root: FleetRoot,
mirror_jobs: Arc<Semaphore>,
states: TaskStates,
}
impl Node {
async fn run(&self, shutdown: CancellationToken) {
info!(
node_id = %self.options.node_id,
root = %self.root.url(),
heartbeat = %self.object_name,
"sleet node starting"
);
let mut poller = ConfigPoller::default();
let mut state = poller.poll(&self.root).await;
log_warnings(&state);
let mut last_poll = Instant::now();
let mut tasks: HashMap<Assignment, RunningTask> = HashMap::new();
loop {
self.put_heartbeat().await;
if last_poll.elapsed() >= state.config.node.config_poll.0 {
state = poller.poll(&self.root).await;
log_warnings(&state);
last_poll = Instant::now();
}
match self.root.list_heartbeats().await {
Ok(entries) => {
self.housekeeping(&entries, &state).await;
let owned = self.owned_assignments(&entries, &state);
self.reconcile(&mut tasks, owned, &state);
}
Err(e) => warn!("failed to LIST nodes/: {e}; keeping current assignments"),
}
tokio::select! {
_ = shutdown.cancelled() => break,
_ = tokio::time::sleep(state.config.node.heartbeat_interval.0) => {}
}
}
info!("shutting down: stopping tasks and deleting heartbeat");
for task in tasks.values() {
task.token.cancel();
}
let stop_all = futures::future::join_all(tasks.into_values().map(|t| t.handle));
let _ = tokio::time::timeout(Duration::from_secs(30), stop_all).await;
if let Err(e) = self
.root
.store()
.delete(&self.root.node_path(&self.object_name))
.await
{
warn!("failed to delete heartbeat on shutdown: {e}");
}
}
async fn put_heartbeat(&self) {
let mut summaries: HashMap<Service, ServiceSummary> = self
.options
.services
.iter()
.map(|&service| (service, ServiceSummary::empty(service)))
.collect();
for ((_, service, _), (_, state)) in self.states.lock().expect("states lock").iter() {
let summary = summaries
.entry(*service)
.or_insert(ServiceSummary::empty(*service));
match state {
TaskState::Running => summary.running += 1,
TaskState::Backoff => summary.backoff += 1,
}
}
let mut services: Vec<ServiceSummary> = summaries.into_values().collect();
services.sort_by_key(|s| s.service.letter());
let body = Heartbeat::new(&self.options.node_id, SLATEDB_VERSION, services);
let json = serde_json::to_vec(&body).expect("heartbeat serializes");
let path = self.root.node_path(&self.object_name);
if let Err(e) = self.root.store().put(&path, PutPayload::from(json)).await {
warn!("failed to PUT heartbeat: {e}");
}
}
async fn housekeeping(&self, entries: &[HeartbeatEntry], state: &FleetState) {
let own = self.root.node_path(&self.object_name);
let long_dead = state.config.node.heartbeat_timeout.0 * 10;
for entry in entries {
let stale_own = entry.node_id == self.options.node_id && entry.location != own;
if (stale_own || entry.age >= long_dead)
&& let Err(e) = self.root.store().delete(&entry.location).await
{
warn!("failed to delete heartbeat {}: {e}", entry.location);
}
}
}
fn owned_assignments(
&self,
entries: &[HeartbeatEntry],
state: &FleetState,
) -> HashMap<Assignment, Arc<ResolvedServices>> {
owned_assignments(
&self.options.node_id,
&self.options.services,
entries,
state,
)
}
fn reconcile(
&self,
tasks: &mut HashMap<Assignment, RunningTask>,
owned: HashMap<Assignment, Arc<ResolvedServices>>,
state: &FleetState,
) {
let stop: Vec<Assignment> = tasks
.iter()
.filter(|(key, task)| match owned.get(*key) {
Some(resolved) => fingerprint(resolved) != task.fingerprint,
None => true,
})
.map(|(key, _)| key.clone())
.collect();
for key in stop {
if let Some(task) = tasks.remove(&key) {
info!(
database = %key.0,
service = key.1.as_str(),
target = key.2.as_deref().unwrap_or(""),
"stopping task"
);
task.token.cancel();
}
}
tasks.retain(|_, task| !task.handle.is_finished());
for (key, resolved) in owned {
if tasks.contains_key(&key) {
continue;
}
info!(
database = %key.0,
service = key.1.as_str(),
target = key.2.as_deref().unwrap_or(""),
"starting task"
);
let token = CancellationToken::new();
let handle = tokio::spawn(supervise(
key.clone(),
resolved.clone(),
self.mirror_jobs.clone(),
self.options.rclone.clone(),
self.states.clone(),
token.clone(),
state.config.node.heartbeat_interval.0,
));
tasks.insert(
key,
RunningTask {
token,
handle,
fingerprint: fingerprint(&resolved),
},
);
}
}
}
async fn supervise(
key: Assignment,
resolved: Arc<ResolvedServices>,
mirror_jobs: Arc<Semaphore>,
rclone: Option<String>,
states: TaskStates,
token: CancellationToken,
heartbeat_interval: Duration,
) {
let (url, service, target) = key.clone();
let run = move |child: CancellationToken| {
let url = url.clone();
let target = target.clone();
let resolved = resolved.clone();
let mirror_jobs = mirror_jobs.clone();
let rclone = rclone.clone();
async move {
if service == Service::Mirror {
let name = target.expect("mirror assignments carry a target");
let Some(applied) = mirror::applied_targets(&url, &resolved.mirror)
.into_iter()
.find(|t| t.name == name)
else {
return Ok(());
};
let source = DatabaseHandle::open(&url)?;
let dest = DatabaseHandle::open(&applied.destination)?;
mirror::run_mirror(&source, &dest, &applied, mirror_jobs, rclone, child)
.await
.map_err(services::ServiceError::from)
} else {
let db = DatabaseHandle::open(&url)?;
services::run_service(&db, service, &resolved, child).await
}
}
};
supervise_with(run, key, states, token, heartbeat_interval).await;
}
async fn supervise_with<F, Fut>(
mut run: F,
key: Assignment,
states: TaskStates,
token: CancellationToken,
heartbeat_interval: Duration,
) where
F: FnMut(CancellationToken) -> Fut,
Fut: Future<Output = Result<(), services::ServiceError>>,
{
let instance = NEXT_INSTANCE.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let (url, service, _) = &key;
let mut backoff = Duration::from_secs(1);
const MAX_BACKOFF: Duration = Duration::from_secs(60);
loop {
set_state(&states, &key, instance, TaskState::Running);
let result = run(token.child_token()).await;
if token.is_cancelled() {
break;
}
let delay = match result {
Ok(()) => break,
Err(e) if e.is_fenced() => {
info!(database = %url, "coordinator fenced; retrying after one heartbeat interval");
backoff = Duration::from_secs(1);
heartbeat_interval
}
Err(e) => {
warn!(database = %url, service = service.as_str(), "task failed: {e}");
backoff = (backoff * 2).min(MAX_BACKOFF);
backoff
}
};
set_state(&states, &key, instance, TaskState::Backoff);
tokio::select! {
_ = token.cancelled() => break,
_ = tokio::time::sleep(delay) => {}
}
}
let mut states = states.lock().expect("states lock");
if states.get(&key).is_some_and(|(id, _)| *id == instance) {
states.remove(&key);
}
}
fn set_state(states: &TaskStates, key: &Assignment, instance: u64, state: TaskState) {
let mut states = states.lock().expect("states lock");
match states.get(key) {
Some((id, _)) if *id > instance => {}
_ => {
states.insert(key.clone(), (instance, state));
}
}
}
fn fingerprint(resolved: &ResolvedServices) -> u64 {
use std::hash::{Hash, Hasher};
let mut hasher = std::hash::DefaultHasher::new();
format!("{resolved:?}").hash(&mut hasher);
hasher.finish()
}
fn log_warnings(state: &FleetState) {
for warning in &state.warnings {
warn!("{warning}");
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::{DatabaseConfig, SleetConfig, WorkersOverrides};
use crate::root::HeartbeatEntry;
use crate::testing::TestStore;
use object_store::path::Path as StorePath;
use slatedb::{CloseReason, Error as SlateError};
fn node(node_id: &str, services: &[Service]) -> Node {
let options = NodeOptions {
node_id: node_id.into(),
services: services.to_vec(),
..NodeOptions::default()
};
Node {
object_name: heartbeat::object_name(&options.node_id, &options.services),
mirror_jobs: Arc::new(Semaphore::new(1)),
states: TaskStates::default(),
root: FleetRoot::from_parts(
TestStore::in_memory(),
StorePath::from("fleet"),
"memory:///fleet",
),
options,
}
}
fn entry(node_id: &str, services: &[Service], age_secs: u64) -> HeartbeatEntry {
HeartbeatEntry {
node_id: node_id.into(),
services: services.to_vec(),
age: Duration::from_secs(age_secs),
location: StorePath::from(format!(
"fleet/nodes/{}",
heartbeat::object_name(node_id, services)
)),
}
}
fn state(databases: &[(&str, DatabaseConfig)]) -> FleetState {
FleetState {
config: SleetConfig::default(),
databases: databases
.iter()
.map(|(url, db)| (url.to_string(), db.clone()))
.collect(),
warnings: vec![],
}
}
#[test]
fn owned_assignments_follow_the_ranking() {
let dbs = state(&[
("s3://b/db1", DatabaseConfig::default()),
("s3://b/db2", DatabaseConfig::default()),
]);
let entries = vec![entry("n1", &Service::ALL, 1), entry("n2", &Service::ALL, 1)];
let n1_owned = node("n1", &Service::ALL).owned_assignments(&entries, &dbs);
let n2_owned = node("n2", &Service::ALL).owned_assignments(&entries, &dbs);
for url in ["s3://b/db1", "s3://b/db2"] {
for service in Service::ALL {
if service == Service::Mirror {
continue;
}
let expected = placement::owners(url, service, 1, &["n1", "n2"])[0];
let key = (url.to_string(), service, None);
assert_eq!(n1_owned.contains_key(&key), expected == "n1");
assert_eq!(n2_owned.contains_key(&key), expected == "n2");
}
}
}
#[test]
fn owned_assignments_respect_roles() {
let dbs = state(&[("s3://b/db1", DatabaseConfig::default())]);
let entries = vec![
entry("n1", &[Service::Gc], 1),
entry("n2", &[Service::CompactionWorkers], 1),
];
let gc_only = node("n1", &[Service::Gc]).owned_assignments(&entries, &dbs);
assert!(gc_only.contains_key(&("s3://b/db1".into(), Service::Gc, None)));
assert!(!gc_only.contains_key(&("s3://b/db1".into(), Service::CompactionWorkers, None)));
assert!(!gc_only.contains_key(&("s3://b/db1".into(), Service::CompactorCoordinator, None)));
let workers_only =
node("n2", &[Service::CompactionWorkers]).owned_assignments(&entries, &dbs);
assert_eq!(workers_only.len(), 1);
assert!(workers_only.contains_key(&(
"s3://b/db1".into(),
Service::CompactionWorkers,
None
)));
}
#[test]
fn owned_assignments_count_disabled_and_dead() {
let two_workers = DatabaseConfig {
compaction_workers: Some(WorkersOverrides {
count: Some(2),
..Default::default()
}),
..Default::default()
};
let disabled = DatabaseConfig {
services: Some(vec![]),
..Default::default()
};
let dbs = state(&[("s3://b/db1", two_workers), ("s3://b/off", disabled)]);
let entries = vec![
entry("n1", &Service::ALL, 1),
entry("n2", &Service::ALL, 1),
entry("n3", &Service::ALL, 999), ];
let worker_key = (
"s3://b/db1".to_string(),
Service::CompactionWorkers,
None::<String>,
);
let n1 = node("n1", &Service::ALL).owned_assignments(&entries, &dbs);
let n2 = node("n2", &Service::ALL).owned_assignments(&entries, &dbs);
let n3 = node("n3", &Service::ALL).owned_assignments(&entries, &dbs);
assert!(n1.contains_key(&worker_key) && n2.contains_key(&worker_key));
for key in n3.keys() {
let (url, service, _) = key;
let count = if *service == Service::CompactionWorkers {
2
} else {
1
};
assert!(
placement::owners(url, *service, count, &["n1", "n2", "n3"]).contains(&"n3"),
"{key:?}"
);
}
for owned in [&n1, &n2] {
assert!(!owned.keys().any(|(url, ..)| url == "s3://b/off"));
}
}
#[test]
fn owned_assignments_place_mirror_targets() {
use crate::config::{MirrorOverrides, MirrorTargetOverrides};
let mirrored = DatabaseConfig {
mirror: Some(MirrorOverrides {
targets: [
(
"dr".to_string(),
MirrorTargetOverrides {
url: Some("s3://dr/db1".into()),
..Default::default()
},
),
(
"backup".to_string(),
MirrorTargetOverrides {
url: Some("gs://backups/db1".into()),
..Default::default()
},
),
]
.into(),
}),
..Default::default()
};
let dbs = state(&[
("s3://b/db1", mirrored),
("s3://b/plain", DatabaseConfig::default()),
]);
let entries = vec![
entry("n1", &Service::ALL, 1),
entry("n2", &Service::ALL, 1),
entry("n3", &[Service::Gc], 1), ];
let mut owners_seen = Vec::new();
for id in ["n1", "n2", "n3"] {
let owned = node(id, &Service::ALL).owned_assignments(&entries, &dbs);
for target in ["dr", "backup"] {
let key = (
"s3://b/db1".to_string(),
Service::Mirror,
Some(target.to_string()),
);
let expected = placement::owner_target("s3://b/db1", target, &["n1", "n2"]);
assert_eq!(
owned.contains_key(&key),
expected == Some(id),
"{id} {target}"
);
if owned.contains_key(&key) {
owners_seen.push((id, target));
}
}
assert!(
!owned
.keys()
.any(|(url, service, _)| url == "s3://b/plain" && *service == Service::Mirror)
);
}
assert_eq!(owners_seen.len(), 2, "each target owned exactly once");
}
#[tokio::test(flavor = "multi_thread")]
async fn reconcile_diffs_tasks_against_ownership() {
let node = node("n1", &Service::ALL);
let fleet = state(&[]);
let mut tasks = HashMap::new();
let key = ("memory:///db1".to_string(), Service::Gc, None);
let resolved = Arc::new(SleetConfig::default().resolve(None));
let mut owned = HashMap::new();
owned.insert(key.clone(), resolved.clone());
node.reconcile(&mut tasks, owned.clone(), &fleet);
assert!(tasks.contains_key(&key));
let first_token = tasks[&key].token.clone();
node.reconcile(&mut tasks, owned.clone(), &fleet);
assert!(!first_token.is_cancelled());
let mut changed = SleetConfig::default().resolve(None);
changed.workers.count = 7;
let mut owned_changed = HashMap::new();
owned_changed.insert(key.clone(), Arc::new(changed));
node.reconcile(&mut tasks, owned_changed, &fleet);
assert!(first_token.is_cancelled());
assert!(tasks.contains_key(&key));
let second_token = tasks[&key].token.clone();
node.reconcile(&mut tasks, HashMap::new(), &fleet);
assert!(second_token.is_cancelled());
assert!(tasks.is_empty());
}
fn fenced() -> services::ServiceError {
services::ServiceError::SlateDb(SlateError::closed(
"fenced by another coordinator".into(),
CloseReason::Fenced,
))
}
fn plain() -> services::ServiceError {
services::ServiceError::SlateDb(SlateError::unavailable("boom".into()))
}
#[tokio::test(start_paused = true)]
async fn supervisor_backoff_policy() {
let outcomes = Arc::new(Mutex::new(vec![
Err(fenced()),
Err(plain()),
Err(plain()),
Ok(()),
]));
let calls = Arc::new(Mutex::new(Vec::new()));
let run = {
let outcomes = outcomes.clone();
let calls = calls.clone();
move |_child: CancellationToken| {
let outcomes = outcomes.clone();
let calls = calls.clone();
async move {
calls.lock().unwrap().push(tokio::time::Instant::now());
outcomes.lock().unwrap().remove(0)
}
}
};
let heartbeat_interval = Duration::from_secs(7);
supervise_with(
run,
("db".into(), Service::CompactorCoordinator, None),
TaskStates::default(),
CancellationToken::new(),
heartbeat_interval,
)
.await;
let calls = calls.lock().unwrap();
assert_eq!(calls.len(), 4);
assert_eq!(calls[1] - calls[0], heartbeat_interval);
assert_eq!(calls[2] - calls[1], Duration::from_secs(2));
assert_eq!(calls[3] - calls[2], Duration::from_secs(4));
}
#[tokio::test(start_paused = true)]
async fn supervisor_exits_on_cancel() {
let states = TaskStates::default();
let token = CancellationToken::new();
let run = |_child: CancellationToken| async { Err(plain()) };
let handle = tokio::spawn(supervise_with(
run,
("db".into(), Service::Gc, None),
states.clone(),
token.clone(),
Duration::from_secs(10),
));
tokio::time::sleep(Duration::from_millis(100)).await;
token.cancel();
handle.await.unwrap();
assert!(states.lock().unwrap().is_empty());
}
}