use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering as AtomicOrdering};
use std::time::Duration;
use anyhow::{Context, Result};
use tokio::process::{Child, Command};
use tokio::sync::{RwLock, watch};
use super::{EmbedderClient, StdioEmbedderClient};
#[derive(Debug, Clone)]
pub struct SupervisorConfig {
pub max_restarts: u32,
pub backoff_max_secs: u64,
pub startup_timeout_secs: u64,
pub sidecar_batch_size: Option<usize>,
pub wedge_reset_secs: u64,
}
impl Default for SupervisorConfig {
fn default() -> Self {
Self {
max_restarts: 5,
backoff_max_secs: 60,
startup_timeout_secs: 5,
sidecar_batch_size: None,
wedge_reset_secs: 300,
}
}
}
impl SupervisorConfig {
pub fn from_env() -> Self {
let def = Self::default();
Self {
max_restarts: parse_env("TRUSTY_EMBEDDERD_MAX_RESTARTS", def.max_restarts),
backoff_max_secs: parse_env(
"TRUSTY_EMBEDDERD_RESTART_BACKOFF_MAX_SECS",
def.backoff_max_secs,
),
startup_timeout_secs: parse_env(
"TRUSTY_EMBEDDERD_STARTUP_TIMEOUT_SECS",
def.startup_timeout_secs,
),
sidecar_batch_size: None,
wedge_reset_secs: parse_env("TRUSTY_EMBEDDERD_WEDGE_RESET_SECS", def.wedge_reset_secs),
}
}
}
pub const DEFAULT_CUDA_SIDECAR_BATCH_CAP: usize = 64;
pub fn cuda_sidecar_batch_cap() -> usize {
static CACHED: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*CACHED.get_or_init(|| {
std::env::var("TRUSTY_CUDA_SIDECAR_BATCH_CAP")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or(DEFAULT_CUDA_SIDECAR_BATCH_CAP)
.clamp(1, 512)
})
}
pub fn sidecar_batch_size(
resolved: usize,
is_coreml: bool,
coreml_cap: usize,
is_cuda: bool,
cuda_cap: usize,
) -> usize {
let raw = if is_coreml {
if coreml_cap == 0 {
tracing::warn!(
resolved,
"sidecar_batch_size: CoreML batch cap resolved to 0 — likely a \
resolve_coreml_batch_size() misconfiguration. Clamping to 1, \
which will be very slow (one embedding per ONNX call). \
Check TRUSTY_COREML_TRIPWIRE_MB and available system RAM."
);
}
resolved.min(coreml_cap)
} else if is_cuda {
if cuda_cap == 0 {
tracing::warn!(
resolved,
"sidecar_batch_size: CUDA batch cap resolved to 0 — likely a \
misconfiguration. Clamping to 1. \
Check TRUSTY_CUDA_SIDECAR_BATCH_CAP."
);
}
resolved.min(cuda_cap)
} else {
resolved
};
raw.max(1)
}
fn parse_env<T: std::str::FromStr + Copy>(name: &str, default: T) -> T {
std::env::var(name)
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(default)
}
pub struct EmbedderSupervisor {
binary_path: PathBuf,
child: Arc<tokio::sync::Mutex<Option<Child>>>,
client_slot: Arc<RwLock<Arc<dyn EmbedderClient>>>,
child_pid_slot: Arc<AtomicU32>,
unhealthy_signal: watch::Receiver<bool>,
config: SupervisorConfig,
}
impl EmbedderSupervisor {
pub async fn spawn_stdio(
binary_path: impl Into<PathBuf>,
config: SupervisorConfig,
) -> Result<(Self, Arc<RwLock<Arc<dyn EmbedderClient>>>, Arc<AtomicU32>)> {
let binary_path = binary_path.into();
let (child, client) = spawn_child(&binary_path, &config).await?;
let unhealthy_signal = client.unhealthy_signal();
let initial_pid: u32 = child.id().unwrap_or(0);
let child_pid_slot = Arc::new(AtomicU32::new(initial_pid));
let client_slot: Arc<RwLock<Arc<dyn EmbedderClient>>> =
Arc::new(RwLock::new(Arc::new(client)));
let client_slot_clone = Arc::clone(&client_slot);
let child_pid_slot_clone = Arc::clone(&child_pid_slot);
let supervisor = Self {
binary_path,
child: Arc::new(tokio::sync::Mutex::new(Some(child))),
client_slot,
child_pid_slot,
unhealthy_signal,
config,
};
Ok((supervisor, client_slot_clone, child_pid_slot_clone))
}
pub fn start_supervisor_task(self) -> SupervisorHandle {
let (shutdown_tx, shutdown_rx) = watch::channel(false);
let join = tokio::spawn(supervision_loop(
self.binary_path,
self.child,
self.client_slot,
self.child_pid_slot,
self.unhealthy_signal,
self.config,
shutdown_rx,
));
SupervisorHandle { shutdown_tx, join }
}
pub async fn shutdown(self) {
let mut guard = self.child.lock().await;
if let Some(mut child) = guard.take() {
let _ = child.kill().await;
let _ = child.wait().await;
tracing::info!("EmbedderSupervisor: sidecar terminated on shutdown");
}
}
}
#[must_use = "dropping the handle without calling shutdown() leaves the \
sidecar running in the background with no way to stop it \
cooperatively later"]
pub struct SupervisorHandle {
shutdown_tx: watch::Sender<bool>,
join: tokio::task::JoinHandle<()>,
}
impl SupervisorHandle {
pub async fn shutdown(self) {
let _ = self.shutdown_tx.send(true);
if let Err(e) = self.join.await {
tracing::warn!("EmbedderSupervisor: supervision task join error on shutdown: {e}");
}
}
}
async fn spawn_child(
binary_path: &Path,
config: &SupervisorConfig,
) -> Result<(Child, StdioEmbedderClient)> {
use std::process::Stdio;
let mut cmd = Command::new(binary_path);
cmd.arg("--stdio")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::inherit())
.kill_on_drop(true);
if let Some(bs) = config.sidecar_batch_size {
cmd.env("TRUSTY_EMBED_BATCH_SIZE", bs.to_string());
tracing::debug!(
bs,
"EmbedderSupervisor: forwarding TRUSTY_EMBED_BATCH_SIZE={bs}"
);
}
let mut child = cmd.spawn().with_context(|| {
format!(
"spawn trusty-embedderd --stdio from {}",
binary_path.display()
)
})?;
let stdin = child
.stdin
.take()
.context("child stdin handle missing (expected Stdio::piped)")?;
let stdout = child
.stdout
.take()
.context("child stdout handle missing (expected Stdio::piped)")?;
let client = StdioEmbedderClient::new(stdin, stdout);
let probe_result = tokio::time::timeout(
Duration::from_secs(config.startup_timeout_secs),
client.embed_batch(vec!["trusty-embedderd startup probe".to_string()]),
)
.await;
match probe_result {
Ok(Ok(_)) => {
tracing::info!(
binary = %binary_path.display(),
"EmbedderSupervisor: sidecar started and responding"
);
}
Ok(Err(e)) => {
anyhow::bail!("sidecar startup probe failed: {e}");
}
Err(_elapsed) => {
anyhow::bail!(
"sidecar did not respond within {}s (TRUSTY_EMBEDDERD_STARTUP_TIMEOUT_SECS={})",
config.startup_timeout_secs,
config.startup_timeout_secs
);
}
}
Ok((child, client))
}
enum RestartTrigger {
ProcessExit(std::process::ExitStatus),
Unhealthy,
}
fn wedge_counter_should_reset(
elapsed_since_last_wedge: Option<Duration>,
wedge_reset_secs: u64,
) -> bool {
match elapsed_since_last_wedge {
Some(elapsed) => elapsed >= Duration::from_secs(wedge_reset_secs),
None => false,
}
}
fn should_give_up(
consecutive_failures: u32,
consecutive_wedge_restarts: u32,
max_restarts: u32,
) -> bool {
consecutive_failures > max_restarts || consecutive_wedge_restarts > max_restarts
}
#[cfg(test)]
pub(crate) static SUPERVISION_LOOP_ITERATIONS: std::sync::atomic::AtomicU64 =
std::sync::atomic::AtomicU64::new(0);
async fn supervision_loop(
binary_path: PathBuf,
child_slot: Arc<tokio::sync::Mutex<Option<Child>>>,
client_slot: Arc<RwLock<Arc<dyn EmbedderClient>>>,
child_pid_slot: Arc<AtomicU32>,
mut unhealthy_signal: watch::Receiver<bool>,
config: SupervisorConfig,
mut shutdown_rx: watch::Receiver<bool>,
) {
let mut consecutive_failures: u32 = 0;
let mut consecutive_wedge_restarts: u32 = 0;
let mut last_wedge_restart_at: Option<tokio::time::Instant> = None;
let mut shutdown_closed = false;
loop {
#[cfg(test)]
SUPERVISION_LOOP_ITERATIONS.fetch_add(1, AtomicOrdering::Relaxed);
if wedge_counter_should_reset(
last_wedge_restart_at.map(|t| t.elapsed()),
config.wedge_reset_secs,
) {
tracing::info!(
"EmbedderSupervisor: {}s without a further wedge — resetting wedge-restart \
escalation counter (was {consecutive_wedge_restarts})",
config.wedge_reset_secs,
);
consecutive_wedge_restarts = 0;
last_wedge_restart_at = None;
}
let trigger = {
let mut guard = child_slot.lock().await;
match guard.as_mut() {
Some(child) => {
tokio::select! {
wait_result = child.wait() => match wait_result {
Ok(status) => RestartTrigger::ProcessExit(status),
Err(e) => {
tracing::error!("EmbedderSupervisor: wait() failed: {e}");
child_pid_slot.store(0, AtomicOrdering::Release);
return;
}
},
changed = unhealthy_signal.changed() => {
if changed.is_err() {
tracing::debug!(
"EmbedderSupervisor: unhealthy_signal channel closed \
without firing — continuing to watch child.wait()"
);
continue;
}
RestartTrigger::Unhealthy
}
changed_result = shutdown_rx.changed(), if !shutdown_closed => {
let Ok(()) = changed_result else {
shutdown_closed = true;
continue;
};
if !*shutdown_rx.borrow() {
continue;
}
tracing::info!(
"EmbedderSupervisor: shutdown requested — stopping the \
sidecar and supervision loop (no respawn)"
);
if let Err(e) = child.start_kill() {
tracing::warn!(
"EmbedderSupervisor: start_kill on shutdown failed \
(may have already exited): {e}"
);
}
let _ = child.wait().await;
child_pid_slot.store(0, AtomicOrdering::Release);
return;
}
}
}
None => {
child_pid_slot.store(0, AtomicOrdering::Release);
return;
}
}
};
let is_wedge_restart = matches!(trigger, RestartTrigger::Unhealthy);
match trigger {
RestartTrigger::ProcessExit(exit_status) => {
child_pid_slot.store(0, AtomicOrdering::Release);
if exit_status.success() {
tracing::info!(
"EmbedderSupervisor: sidecar exited cleanly — stopping supervision"
);
return;
}
tracing::warn!(
"EmbedderSupervisor: sidecar exited with {:?} (failure #{}/{})",
exit_status.code(),
consecutive_failures + 1,
config.max_restarts,
);
}
RestartTrigger::Unhealthy => {
consecutive_wedge_restarts += 1;
last_wedge_restart_at = Some(tokio::time::Instant::now());
tracing::warn!(
"EmbedderSupervisor: sidecar client reported unhealthy (reader task \
died, or accumulating call timeouts indicate a wedged process) — \
forcing restart (failure #{}/{}, wedge-restart #{}/{})",
consecutive_failures + 1,
config.max_restarts,
consecutive_wedge_restarts,
config.max_restarts,
);
let mut guard = child_slot.lock().await;
if let Some(mut child) = guard.take() {
if let Err(e) = child.start_kill() {
tracing::warn!(
"EmbedderSupervisor: start_kill on wedged sidecar failed \
(may have already exited): {e}"
);
}
let _ = child.wait().await;
}
drop(guard);
child_pid_slot.store(0, AtomicOrdering::Release);
}
}
consecutive_failures += 1;
if should_give_up(
consecutive_failures,
consecutive_wedge_restarts,
config.max_restarts,
) {
tracing::error!(
"EmbedderSupervisor: exceeded max_restarts={} ({}) — giving up. \
Set TRUSTY_EMBEDDERD_MAX_RESTARTS to increase the limit, or \
TRUSTY_EMBEDDERD_WEDGE_RESET_SECS if wedges are recurring faster \
than the sustained-health reset window.",
config.max_restarts,
if consecutive_wedge_restarts > config.max_restarts {
"wedge-restart storm — recurring despite successful respawns"
} else {
"process-exit crash storm"
},
);
return;
}
let backoff_attempt = if is_wedge_restart {
consecutive_wedge_restarts
} else {
consecutive_failures
};
let delay_secs = (1u64 << backoff_attempt.min(16)).min(config.backoff_max_secs);
tracing::info!(
"EmbedderSupervisor: restarting sidecar in {delay_secs}s (attempt \
{consecutive_failures}{})",
if is_wedge_restart {
format!(", wedge-triggered, wedge-restart #{consecutive_wedge_restarts}")
} else {
String::new()
},
);
tokio::select! {
_ = tokio::time::sleep(Duration::from_secs(delay_secs)) => {}
changed_result = shutdown_rx.changed(), if !shutdown_closed => {
match changed_result {
Err(_) => {
shutdown_closed = true;
tokio::time::sleep(Duration::from_secs(delay_secs)).await;
}
Ok(()) if *shutdown_rx.borrow() => {
tracing::info!(
"EmbedderSupervisor: shutdown requested during respawn \
back-off — stopping without respawn"
);
child_pid_slot.store(0, AtomicOrdering::Release);
return;
}
Ok(()) => {}
}
}
}
match spawn_child(&binary_path, &config).await {
Ok((new_child, new_client)) => {
let new_pid = new_child.id().unwrap_or(0);
unhealthy_signal = new_client.unhealthy_signal();
{
let mut client_guard = client_slot.write().await;
*client_guard = Arc::new(new_client);
}
{
let mut child_guard = child_slot.lock().await;
*child_guard = Some(new_child);
}
child_pid_slot.store(new_pid, AtomicOrdering::Release);
consecutive_failures = 0;
tracing::info!(
"EmbedderSupervisor: sidecar restarted successfully (pid={new_pid})"
);
}
Err(e) => {
tracing::error!("EmbedderSupervisor: respawn failed: {e:#}");
}
}
}
}
pub fn locate_embedderd_binary() -> Result<PathBuf> {
if let Ok(explicit) = std::env::var("TRUSTY_EMBEDDERD_BIN") {
let p = PathBuf::from(&explicit);
if p.is_file() {
return Ok(p);
}
anyhow::bail!("TRUSTY_EMBEDDERD_BIN={explicit:?} does not point to an existing file");
}
if let Ok(exe) = std::env::current_exe()
&& let Some(dir) = exe.parent()
{
let sibling = dir.join("trusty-embedderd");
if sibling.is_file() {
return Ok(sibling);
}
let sibling_exe = dir.join("trusty-embedderd.exe");
if sibling_exe.is_file() {
return Ok(sibling_exe);
}
}
if let Ok(path) = which_embedderd() {
return Ok(path);
}
anyhow::bail!(
"could not locate trusty-embedderd binary. \
Set TRUSTY_EMBEDDERD_BIN=/path/to/trusty-embedderd or ensure it is on PATH."
)
}
fn which_embedderd() -> Result<PathBuf> {
let path_var = std::env::var("PATH").unwrap_or_default();
let sep = if cfg!(windows) { ';' } else { ':' };
for dir in path_var.split(sep) {
let candidate = PathBuf::from(dir).join("trusty-embedderd");
if candidate.is_file() {
return Ok(candidate);
}
#[cfg(windows)]
{
let candidate_exe = PathBuf::from(dir).join("trusty-embedderd.exe");
if candidate_exe.is_file() {
return Ok(candidate_exe);
}
}
}
anyhow::bail!("trusty-embedderd not found on PATH")
}
#[cfg(test)]
#[path = "supervisor_tests.rs"]
mod tests;