use std::sync::Arc;
use std::time::Duration;
use dashmap::DashMap;
use tokio_util::sync::CancellationToken;
use cf_system_sdks::directory::{DirectoryClient, RegisterInstanceInfo};
use super::readiness::ReadinessState;
const INITIAL_BACKOFF: Duration = Duration::from_millis(100);
const MAX_BACKOFF: Duration = Duration::from_secs(30);
const RE_REGISTER_INTERVAL: Duration = Duration::from_secs(30);
fn next_backoff(current: Duration) -> Duration {
current.saturating_mul(2).min(MAX_BACKOFF)
}
async fn sleep_or_cancel(dur: Duration, cancel: &CancellationToken) -> bool {
tokio::select! {
() = cancel.cancelled() => false,
() = tokio::time::sleep(dur) => true,
}
}
#[derive(Debug, Default)]
pub struct ResolvedRestEndpoints {
inner: DashMap<String, String>,
}
impl ResolvedRestEndpoints {
#[must_use]
pub fn new() -> Self {
Self::default()
}
pub fn set(&self, gear: impl Into<String>, uri: impl Into<String>) {
self.inner.insert(gear.into(), uri.into());
}
#[must_use]
pub fn get(&self, gear: &str) -> Option<String> {
self.inner.get(gear).map(|v| v.value().clone())
}
#[must_use]
pub fn len(&self) -> usize {
self.inner.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.inner.is_empty()
}
}
async fn register_once_with_backoff(
directory: &Arc<dyn DirectoryClient>,
info: &RegisterInstanceInfo,
cancel: &CancellationToken,
) -> bool {
let mut backoff = INITIAL_BACKOFF;
loop {
if cancel.is_cancelled() {
return false;
}
match directory.register_instance(info.clone()).await {
Ok(()) => {
tracing::info!(gear = %info.gear, instance = %info.instance_id, "registered with DirectoryService");
return true;
}
Err(e) => {
tracing::warn!(
gear = %info.gear,
error = %e,
backoff_ms = backoff.as_millis(),
"registration attempt failed; retrying"
);
if !sleep_or_cancel(backoff, cancel).await {
return false;
}
backoff = next_backoff(backoff);
}
}
}
}
pub(super) async fn presence_loop(
directory: Arc<dyn DirectoryClient>,
info: RegisterInstanceInfo,
heartbeat_interval: Duration,
cancel: CancellationToken,
) {
if !register_once_with_backoff(&directory, &info, &cancel).await {
return;
}
let heartbeat_interval = heartbeat_interval.max(Duration::from_secs(1));
let mut heartbeat = tokio::time::interval(heartbeat_interval);
heartbeat.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
heartbeat.tick().await;
let mut reregister = tokio::time::interval(RE_REGISTER_INTERVAL);
reregister.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
reregister.tick().await;
loop {
tokio::select! {
() = cancel.cancelled() => return,
_ = heartbeat.tick() => {
if let Err(e) = directory.send_heartbeat(&info.gear, &info.instance_id).await {
tracing::warn!(
gear = %info.gear,
instance = %info.instance_id,
error = %e,
"heartbeat failed; re-registering to self-heal"
);
if !register_once_with_backoff(&directory, &info, &cancel).await {
return;
}
} else {
tracing::trace!(gear = %info.gear, "heartbeat sent");
}
}
_ = reregister.tick() => {
if !register_once_with_backoff(&directory, &info, &cancel).await {
return;
}
}
}
}
}
async fn resolve_one_dep(
directory: Arc<dyn DirectoryClient>,
dep: String,
readiness: Arc<ReadinessState>,
resolved: Arc<ResolvedRestEndpoints>,
cancel: CancellationToken,
) {
let mut backoff = INITIAL_BACKOFF;
loop {
if cancel.is_cancelled() {
return;
}
match directory.resolve_rest_service(&dep).await {
Ok(endpoint) => {
tracing::info!(dep = %dep, endpoint = %endpoint.uri, "resolved REST dependency");
resolved.set(dep.clone(), endpoint.uri);
readiness.mark_dep_resolved(&dep);
return;
}
Err(e) => {
tracing::debug!(dep = %dep, error = %e, "dependency not yet resolvable; retrying");
if !sleep_or_cancel(backoff, &cancel).await {
return;
}
backoff = next_backoff(backoff);
}
}
}
}
pub(super) fn resolve_deps(
directory: &Arc<dyn DirectoryClient>,
deps: Vec<String>,
readiness: &Arc<ReadinessState>,
resolved: &Arc<ResolvedRestEndpoints>,
cancel: &CancellationToken,
) {
for dep in deps {
tokio::spawn(resolve_one_dep(
Arc::clone(directory),
dep,
Arc::clone(readiness),
Arc::clone(resolved),
cancel.clone(),
));
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
#[path = "oop_registration_tests.rs"]
mod tests;