use std::sync::Arc;
mod store_guard;
mod store_identity;
#[cfg(unix)]
use store_guard::ensure_claimed_parent_identity;
#[cfg(unix)]
pub use store_guard::{acquire_daemon_store_guards, bind_daemon_store_files, claim_stores};
pub use store_guard::{assert_daemon_store_identities, DaemonStoreGuard};
#[cfg(unix)]
pub use store_identity::claimed_daemon_store_identity;
mod supervisor_marker;
#[cfg(unix)]
pub use supervisor_marker::supervisor_marker_path;
#[cfg(unix)]
use std::io::Write as _;
#[cfg(unix)]
use std::os::unix::fs::{MetadataExt, PermissionsExt};
#[cfg(unix)]
use std::os::unix::io::AsRawFd;
use std::path::PathBuf;
#[cfg(unix)]
use async_trait::async_trait;
#[cfg(unix)]
use libc;
use serde::{Deserialize, Serialize};
#[cfg(unix)]
use tokio::io::{AsyncReadExt, AsyncWriteExt};
#[cfg(unix)]
use tokio::net::{UnixListener, UnixStream};
#[cfg(unix)]
use crate::pack::RequestIdentity;
#[cfg(unix)]
use khive_db::{run_checkpoint_task, CheckpointConfig, CheckpointLifecycleOwner, ConnectionPool};
mod load_limits;
#[cfg(unix)]
use load_limits::{admit_or_refuse_busy, ConnectionAdmission};
pub use load_limits::{
recall_ledger_snapshot, track_recall_ledger_task, ConnectionCapSnapshot, RecallLedgerSnapshot,
};
pub const MAX_FRAME_BYTES: usize = 8 * 1024 * 1024;
pub const PROTOCOL_VERSION: u32 = 8;
pub const DEFAULT_DEMAND_IDLE_SECS: u64 = 1_800;
#[derive(Serialize, Deserialize, Debug, Clone, Copy, Default, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum DaemonLifetime {
Demand,
#[default]
Persistent,
}
#[derive(Debug, Clone, Copy)]
pub struct DaemonOptions {
pub lifetime: DaemonLifetime,
pub idle_interval: std::time::Duration,
}
impl Default for DaemonOptions {
fn default() -> Self {
Self {
lifetime: DaemonLifetime::Persistent,
idle_interval: std::time::Duration::from_secs(DEFAULT_DEMAND_IDLE_SECS),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct DaemonStartupReport {
pub skipped_components: Vec<String>,
pub idle_ineligible_reasons: Vec<String>,
}
#[derive(Serialize, Deserialize, Debug, Clone, Copy, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum DaemonLifecyclePhase {
Serving,
Draining,
Stopped,
}
#[derive(Serialize, Deserialize, Debug, Clone, Copy, PartialEq, Eq)]
#[serde(rename_all = "snake_case")]
pub enum DaemonShutdownReason {
Idle,
Signal,
}
#[derive(Serialize, Deserialize, Debug, Clone, PartialEq, Eq)]
pub struct DaemonLifecycleSnapshot {
pub lifetime: DaemonLifetime,
pub instance_generation: String,
pub effective_idle_interval_ms: u64,
pub phase: DaemonLifecyclePhase,
pub shutdown_reason: Option<DaemonShutdownReason>,
pub skipped_components: Vec<String>,
pub idle_ineligible_reasons: Vec<String>,
pub ordinary_requests: usize,
pub idle_blockers: Vec<String>,
}
#[cfg(unix)]
struct DaemonLifecycle {
options: DaemonOptions,
state: std::sync::Mutex<DaemonLifecycleState>,
connections: ConnectionAdmission,
}
#[cfg(unix)]
struct DaemonLifecycleState {
snapshot: DaemonLifecycleSnapshot,
last_request_completion: Option<tokio::time::Instant>,
}
#[cfg(unix)]
impl DaemonLifecycle {
fn new(options: DaemonOptions, report: DaemonStartupReport) -> Self {
Self {
options,
state: std::sync::Mutex::new(DaemonLifecycleState {
snapshot: DaemonLifecycleSnapshot {
lifetime: options.lifetime,
instance_generation: uuid::Uuid::new_v4().to_string(),
effective_idle_interval_ms: options
.idle_interval
.as_millis()
.min(u128::from(u64::MAX))
as u64,
phase: DaemonLifecyclePhase::Serving,
shutdown_reason: None,
skipped_components: report.skipped_components,
idle_ineligible_reasons: report.idle_ineligible_reasons,
ordinary_requests: 0,
idle_blockers: Vec::new(),
},
last_request_completion: None,
}),
connections: ConnectionAdmission::from_env(),
}
}
fn snapshot(&self) -> DaemonLifecycleSnapshot {
self.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.snapshot
.clone()
}
fn ready(&self) {
self.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.last_request_completion = Some(tokio::time::Instant::now());
}
fn admit(self: &Arc<Self>) -> Option<OrdinaryRequestGuard> {
let mut state = self
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if state.snapshot.phase != DaemonLifecyclePhase::Serving {
return None;
}
state.snapshot.ordinary_requests += 1;
Some(OrdinaryRequestGuard(Arc::clone(self)))
}
fn try_idle(&self, blockers: impl FnOnce() -> Vec<String>) -> bool {
let mut state = self
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if self.options.lifetime != DaemonLifetime::Demand
|| state.snapshot.phase != DaemonLifecyclePhase::Serving
|| state.snapshot.ordinary_requests != 0
|| !state.snapshot.idle_ineligible_reasons.is_empty()
|| state
.last_request_completion
.is_none_or(|last| last.elapsed() < self.options.idle_interval)
{
return false;
}
state.snapshot.idle_blockers = blockers();
if !state.snapshot.idle_blockers.is_empty() {
return false;
}
state.snapshot.phase = DaemonLifecyclePhase::Draining;
state.snapshot.shutdown_reason = Some(DaemonShutdownReason::Idle);
true
}
fn draining(&self, reason: DaemonShutdownReason) {
let mut state = self
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if state.snapshot.phase != DaemonLifecyclePhase::Stopped {
state.snapshot.phase = DaemonLifecyclePhase::Draining;
state.snapshot.shutdown_reason = Some(reason);
}
}
fn stopped(&self) {
self.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.snapshot
.phase = DaemonLifecyclePhase::Stopped;
}
}
#[cfg(unix)]
struct OrdinaryRequestGuard(Arc<DaemonLifecycle>);
#[cfg(unix)]
impl Drop for OrdinaryRequestGuard {
fn drop(&mut self) {
let mut state = self
.0
.state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
state.snapshot.ordinary_requests -= 1;
state.last_request_completion = Some(tokio::time::Instant::now());
}
}
#[doc(hidden)]
pub const DAEMON_LEXICAL_TIMEOUT_MARKER: &str = "__khive_daemon_lexical_timeout";
const DEFAULT_DRAIN_TIMEOUT_SECS: u64 = 10;
#[cfg(unix)]
const INITIAL_FRAME_READ_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(30);
#[cfg(unix)]
fn next_accept_error_backoff(previous: Option<std::time::Duration>) -> std::time::Duration {
previous
.map(|delay| delay.saturating_mul(2))
.unwrap_or_else(|| std::time::Duration::from_millis(10))
.min(std::time::Duration::from_secs(1))
}
fn khive_dir() -> PathBuf {
khive_root_from(
std::env::var("HOME").ok(),
std::env::var("USERPROFILE").ok(),
)
}
fn khive_root_from(home: Option<String>, userprofile: Option<String>) -> PathBuf {
home.filter(|v| !v.trim().is_empty())
.or_else(|| userprofile.filter(|v| !v.trim().is_empty()))
.map(PathBuf::from)
.unwrap_or_else(last_resort_root)
.join(".khive")
}
pub fn volume_lock_dir() -> Result<PathBuf, khive_db::SqliteError> {
khive_db::default_volume_lock_dir()
}
#[cfg(unix)]
fn last_resort_root() -> PathBuf {
PathBuf::from(".")
}
#[cfg(not(unix))]
fn last_resort_root() -> PathBuf {
std::env::temp_dir()
}
#[cfg(unix)]
const SOCKET_PATH_ENV: &str = "KHIVE_SOCKET";
#[cfg(unix)]
const PID_PATH_ENV: &str = "KHIVE_PID";
#[cfg(unix)]
fn path_override(key: &str) -> Option<PathBuf> {
match std::env::var(key) {
Ok(p) if !p.is_empty() => Some(PathBuf::from(p)),
_ => None,
}
}
#[cfg(unix)]
fn default_socket_path() -> PathBuf {
khive_dir().join("khived.sock")
}
#[cfg(unix)]
fn default_pid_path() -> PathBuf {
khive_dir().join("khived.pid")
}
#[cfg(unix)]
pub fn socket_path() -> PathBuf {
path_override(SOCKET_PATH_ENV).unwrap_or_else(default_socket_path)
}
#[cfg(unix)]
pub fn pid_path() -> PathBuf {
path_override(PID_PATH_ENV).unwrap_or_else(default_pid_path)
}
#[cfg(unix)]
fn ensure_rendezvous_overrides_paired() -> anyhow::Result<()> {
match (path_override(SOCKET_PATH_ENV), path_override(PID_PATH_ENV)) {
(Some(socket), None) => anyhow::bail!(
"refusing to start: {SOCKET_PATH_ENV} is set to {} but {PID_PATH_ENV} is not set. \
The socket and the PID file are two halves of one daemon rendezvous and must move \
together: with only {SOCKET_PATH_ENV} set, this daemon would bind a private socket \
while claiming the shared PID file at {}, which belongs to the default rendezvous \
served on {}. Set {PID_PATH_ENV} to a private path beside the socket, or unset \
{SOCKET_PATH_ENV} to share the default rendezvous.",
socket.display(),
default_pid_path().display(),
default_socket_path().display(),
),
(None, Some(pid)) => anyhow::bail!(
"refusing to start: {PID_PATH_ENV} is set to {} but {SOCKET_PATH_ENV} is not set. \
The socket and the PID file are two halves of one daemon rendezvous and must move \
together: with only {PID_PATH_ENV} set, this daemon would write a private PID file \
while binding the shared socket at {}, the default rendezvous whose owner is \
recorded in {}. Set {SOCKET_PATH_ENV} to a private path beside the PID file, or \
unset {PID_PATH_ENV} to share the default rendezvous.",
pid.display(),
default_socket_path().display(),
default_pid_path().display(),
),
_ => Ok(()),
}
}
pub fn lock_path() -> PathBuf {
if let Ok(p) = std::env::var("KHIVE_LOCK") {
if !p.is_empty() {
return PathBuf::from(p);
}
}
khive_dir().join("khived.recovery.lock")
}
#[cfg(unix)]
pub fn recoverer_lock_path() -> PathBuf {
if let Ok(p) = std::env::var("KHIVE_RECOVERER_LOCK") {
if !p.is_empty() {
return PathBuf::from(p);
}
}
khive_dir().join("khived.recoverer.lock")
}
#[cfg(unix)]
pub const SUPERVISOR_CLAIM_ENV: &str = "KHIVE_SUPERVISOR_CLAIM";
#[cfg(unix)]
fn read_supervisor_marker_claim() -> Option<(u32, String)> {
use std::io::Read;
use std::os::unix::fs::OpenOptionsExt;
let file = std::fs::OpenOptions::new()
.read(true)
.custom_flags(libc::O_NONBLOCK | libc::O_NOFOLLOW)
.open(supervisor_marker_path())
.ok()?;
if !file.metadata().ok()?.is_file() {
return None;
}
let mut marker = String::new();
file.take(4097).read_to_string(&mut marker).ok()?;
if marker.len() > 4096 {
return None;
}
let mut lines = marker.lines();
if lines.next()?.is_empty() {
return None;
}
let pid = lines.next()?.parse::<u32>().ok().filter(|pid| *pid > 0)?;
lines
.next()?
.parse::<u64>()
.ok()
.filter(|seconds| *seconds > 0)?;
let claim = lines.next()?.to_string();
if lines.next().is_some() {
return None;
}
let parsed = uuid::Uuid::parse_str(&claim).ok()?;
if parsed.get_version() != Some(uuid::Version::Random) || parsed.to_string() != claim {
return None;
}
Some((pid, claim))
}
#[cfg(unix)]
fn current_supervisor_claim() -> Option<String> {
let claim = std::env::var(SUPERVISOR_CLAIM_ENV).ok()?;
let (pid, published_claim) = read_supervisor_marker_claim()?;
(pid == std::process::id() && claim == published_claim).then_some(claim)
}
#[cfg(unix)]
fn open_lock_file(path: &std::path::Path) -> std::io::Result<std::fs::File> {
if let Some(parent) = path.parent() {
let _ = std::fs::create_dir_all(parent);
}
std::fs::OpenOptions::new()
.create(true)
.truncate(false)
.write(true)
.open(path)
}
#[cfg(unix)]
fn acquire_flock_blocking(path: &std::path::Path, label: &str) -> Option<std::fs::File> {
let file = match open_lock_file(path) {
Ok(f) => f,
Err(e) => {
tracing::warn!(error = %e, path = ?path, "cannot open {label} lock file");
return None;
}
};
let rc = unsafe { libc::flock(file.as_raw_fd(), libc::LOCK_EX) };
if rc != 0 {
tracing::warn!("flock LOCK_EX failed on {label} lock");
return None;
}
Some(file)
}
#[cfg(unix)]
pub fn acquire_recovery_lock() -> Option<std::fs::File> {
acquire_flock_blocking(&lock_path(), "recovery")
}
#[cfg(unix)]
fn try_acquire_flock_until(
path: &std::path::Path,
deadline: std::time::Instant,
) -> std::io::Result<Option<std::fs::File>> {
let file = open_lock_file(path)?;
let poll_interval = std::time::Duration::from_millis(10);
loop {
let rc = unsafe { libc::flock(file.as_raw_fd(), libc::LOCK_EX | libc::LOCK_NB) };
if rc == 0 {
return Ok(Some(file));
}
let err = std::io::Error::last_os_error();
if err.raw_os_error() != Some(libc::EWOULDBLOCK) {
return Err(err);
}
let now = std::time::Instant::now();
if now >= deadline {
return Ok(None);
}
std::thread::sleep(poll_interval.min(deadline - now));
}
}
#[cfg(unix)]
pub fn try_acquire_daemon_boot_guard_until(
deadline: std::time::Instant,
) -> std::io::Result<Option<DaemonBootGuard>> {
try_acquire_flock_until(&lock_path(), deadline)
}
#[cfg(unix)]
pub fn try_acquire_recoverer_lock_until(
deadline: std::time::Instant,
) -> std::io::Result<Option<std::fs::File>> {
try_acquire_flock_until(&recoverer_lock_path(), deadline)
}
#[cfg(unix)]
pub type DaemonBootGuard = std::fs::File;
#[cfg(unix)]
pub fn acquire_daemon_boot_guard() -> anyhow::Result<DaemonBootGuard> {
acquire_recovery_lock()
.ok_or_else(|| anyhow::anyhow!("failed to acquire daemon boot/recovery lock"))
}
#[cfg(unix)]
#[derive(Clone, Copy, PartialEq, Eq)]
struct SocketIdentity {
dev: u64,
ino: u64,
}
#[cfg(unix)]
fn socket_identity(path: &std::path::Path) -> Option<SocketIdentity> {
use std::os::unix::fs::MetadataExt;
let meta = std::fs::metadata(path).ok()?;
Some(SocketIdentity {
dev: meta.dev(),
ino: meta.ino(),
})
}
#[cfg(unix)]
pub(crate) fn peer_uid(stream: &UnixStream) -> std::io::Result<u32> {
use std::os::fd::AsRawFd;
let fd = stream.as_raw_fd();
#[cfg(any(target_os = "macos", target_os = "ios", target_vendor = "apple"))]
{
let mut uid: libc::uid_t = 0;
let mut gid: libc::gid_t = 0;
let rc = unsafe { libc::getpeereid(fd, &mut uid, &mut gid) };
if rc != 0 {
return Err(std::io::Error::last_os_error());
}
Ok(uid as u32)
}
#[cfg(target_os = "linux")]
{
let mut cred = libc::ucred {
pid: 0,
uid: 0,
gid: 0,
};
let mut len = std::mem::size_of::<libc::ucred>() as libc::socklen_t;
let rc = unsafe {
libc::getsockopt(
fd,
libc::SOL_SOCKET,
libc::SO_PEERCRED,
(&mut cred as *mut libc::ucred).cast::<libc::c_void>(),
&mut len,
)
};
if rc != 0 {
return Err(std::io::Error::last_os_error());
}
Ok(cred.uid)
}
#[cfg(not(any(
target_os = "linux",
target_os = "macos",
target_os = "ios",
target_vendor = "apple"
)))]
{
let _ = fd;
Err(std::io::Error::new(
std::io::ErrorKind::Unsupported,
"peer-credential capture is not implemented for this platform",
))
}
}
#[cfg(unix)]
pub(crate) fn uid_is_permitted(peer: u32, daemon_euid: u32) -> bool {
peer == daemon_euid
}
mod config_id;
#[cfg(test)]
use config_id::parse_config_id;
pub use config_id::{
config_id_extra_embedder_exclusions, config_ids_compatible, first_config_mismatch_field,
};
#[derive(Serialize, Deserialize, Default)]
pub struct DaemonRequestFrame {
pub ops: String,
#[serde(default)]
pub plan: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub presentation: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub presentation_per_op: Option<Vec<Option<String>>>,
pub namespace: String,
#[serde(default)]
pub actor_id: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub process_ref: Option<String>,
#[serde(default)]
pub visible_namespaces: Vec<String>,
#[serde(default)]
pub config_id: String,
#[serde(default)]
pub protocol_version: u32,
#[serde(default)]
pub probe_only: bool,
#[serde(default)]
pub metrics_only: bool,
#[serde(default)]
#[serde(skip_serializing_if = "Option::is_none")]
pub format: Option<String>,
#[serde(default)]
#[serde(skip_serializing_if = "Option::is_none")]
pub format_per_op: Option<Vec<Option<String>>>,
#[serde(default)]
pub from_wire: bool,
#[serde(default)]
#[serde(skip_serializing_if = "Option::is_none")]
pub request_id: Option<u64>,
}
#[derive(Debug, Clone)]
pub struct DaemonDispatchError {
pub message: String,
pub error_detail: serde_json::Value,
}
pub const ERROR_DETAIL_NESTING_DEPTH_LIMIT: usize = 64;
fn error_detail_value_within_limit(value: &serde_json::Value) -> bool {
let mut pending = vec![(value, 0_usize)];
while let Some((value, depth)) = pending.pop() {
match value {
serde_json::Value::Array(items) if depth < ERROR_DETAIL_NESTING_DEPTH_LIMIT => {
pending.extend(items.iter().map(|child| (child, depth + 1)));
}
serde_json::Value::Object(fields) if depth < ERROR_DETAIL_NESTING_DEPTH_LIMIT => {
pending.extend(fields.values().map(|child| (child, depth + 1)));
}
serde_json::Value::Array(_) | serde_json::Value::Object(_) => return false,
_ => {}
}
}
true
}
fn drop_error_detail_iteratively(value: serde_json::Value) {
let mut pending = vec![value];
while let Some(value) = pending.pop() {
match value {
serde_json::Value::Array(items) => pending.extend(items),
serde_json::Value::Object(fields) => pending.extend(fields.into_values()),
_ => {}
}
}
}
impl DaemonDispatchError {
pub fn new(message: impl Into<String>, error_detail: Option<serde_json::Value>) -> Self {
let message = message.into();
let mut fields = match error_detail {
Some(serde_json::Value::Object(fields)) => fields,
Some(data) => serde_json::Map::from_iter([("data".to_string(), data)]),
None => serde_json::Map::new(),
};
let disposition = match fields
.get("domain_disposition")
.and_then(serde_json::Value::as_str)
{
Some("committed") => crate::DomainDisposition::Committed,
Some("not_committed") => crate::DomainDisposition::NotCommitted,
_ => crate::DomainDisposition::Unknown,
};
if disposition != crate::DomainDisposition::Committed {
if let Some(result) = fields.remove("domain_result") {
drop_error_detail_iteratively(result);
}
}
let rejected: Vec<String> = fields
.iter()
.filter(|(_, value)| !error_detail_value_within_limit(value))
.map(|(name, _)| name.clone())
.collect();
let omitted_result = rejected.iter().any(|name| name == "domain_result");
let omitted_detail = !rejected.is_empty();
for name in rejected {
if let Some(value) = fields.remove(&name) {
drop_error_detail_iteratively(value);
}
}
let mut error_detail = serde_json::Value::Object(fields);
if error_detail["kind"].as_str().is_none() {
error_detail["kind"] = serde_json::json!("internal");
}
if error_detail["message"].as_str().is_none() {
error_detail["message"] = serde_json::json!(message);
}
error_detail["domain_disposition"] = serde_json::json!(disposition.as_str());
if omitted_detail {
error_detail["code"] = serde_json::json!(if omitted_result {
"result_too_deep"
} else {
"error_detail_too_deep"
});
}
Self {
message,
error_detail,
}
}
}
#[derive(Serialize, Deserialize, Debug)]
pub struct DaemonResponseFrame {
pub ok: bool,
pub result: Option<String>,
pub error: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error_detail: Option<serde_json::Value>,
pub namespace_mismatch: bool,
#[serde(default)]
pub config_mismatch: bool,
#[serde(default)]
pub served_config_id: Option<String>,
#[serde(default)]
pub version_mismatch: bool,
#[serde(default)]
pub daemon_protocol_version: u32,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub metrics: Option<MetricsSnapshot>,
#[serde(default)]
pub request_id: Option<u64>,
}
#[cfg(unix)]
fn take_daemon_lexical_timeout_marker(raw: String) -> (String, Option<serde_json::Value>) {
if !raw.contains(DAEMON_LEXICAL_TIMEOUT_MARKER) {
return (raw, None);
}
let Ok(mut value) = serde_json::from_str::<serde_json::Value>(&raw) else {
return (raw, None);
};
let Some(fields) = value.as_object_mut() else {
return (raw, None);
};
if !fields
.get("results")
.is_some_and(serde_json::Value::is_array)
{
return (raw, None);
}
let Some(marker) = fields.remove(DAEMON_LEXICAL_TIMEOUT_MARKER) else {
return (raw, None);
};
let detail =
(marker.as_bool() == Some(true)).then(|| serde_json::json!({"lexical_timeout": true}));
(
serde_json::to_string(&value).expect("serde_json::Value is serializable"),
detail,
)
}
#[derive(Serialize, Deserialize, Debug, Clone, Default, PartialEq)]
#[serde(default)]
pub struct CheckpointStoreMetrics {
pub store_id: String,
pub role: String,
pub database: Option<String>,
#[serde(flatten)]
pub timing: khive_db::checkpoint::CheckpointTiming,
}
#[derive(Serialize, Deserialize, Debug, Clone, Default, PartialEq)]
pub struct MetricsSnapshot {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub lifecycle: Option<DaemonLifecycleSnapshot>,
pub wal_pages: Option<u64>,
#[serde(default)]
pub wal_log_frames: Option<u64>,
#[serde(default)]
pub wal_checkpointed_frames: Option<u64>,
#[serde(default)]
pub wal_pending_frames: Option<u64>,
#[serde(default)]
pub wal_physical_bytes: Option<u64>,
#[serde(default)]
pub wal_observed_at_unix_ms: Option<u64>,
#[serde(default)]
pub wal_checkpoint_stores: Vec<CheckpointStoreMetrics>,
pub wal_truncate_attempts: u64,
pub wal_truncate_consecutive_failures: u64,
#[serde(default)]
pub wal_checkpoint_skipped_ticks: u64,
#[serde(default)]
pub wal_checkpoint_consecutive_skips: u64,
#[serde(default)]
pub wal_checkpoint_last_skip_wal_pages: Option<u64>,
pub oldest_pinned_tx_micros: Option<u64>,
pub oldest_pinned_tx_label: Option<String>,
pub open_tx_count: usize,
pub write_queue_depth: Option<usize>,
pub write_queue_capacity: Option<usize>,
#[serde(default)]
pub write_last_queue_wait_micros: Option<u64>,
#[serde(default)]
pub write_last_transaction_acquire_micros: Option<u64>,
#[serde(default)]
pub write_last_body_micros: Option<u64>,
#[serde(default)]
pub write_last_commit_micros: Option<u64>,
#[serde(default)]
pub write_last_total_micros: Option<u64>,
#[serde(default)]
pub write_last_observed_at_unix_ms: Option<u64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub connections: Option<ConnectionCapSnapshot>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub recall_ledger: Option<RecallLedgerSnapshot>,
}
#[cfg(unix)]
pub async fn read_frame<R>(stream: &mut R) -> std::io::Result<Vec<u8>>
where
R: tokio::io::AsyncRead + Unpin,
{
let mut len_buf = [0u8; 4];
stream.read_exact(&mut len_buf).await?;
let len = u32::from_be_bytes(len_buf) as usize;
if len > MAX_FRAME_BYTES {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!("daemon frame of {len} bytes exceeds {MAX_FRAME_BYTES} cap"),
));
}
let mut buf = vec![0u8; len];
stream.read_exact(&mut buf).await?;
Ok(buf)
}
#[cfg(unix)]
fn initial_frame_timeout_error() -> std::io::Error {
std::io::Error::new(
std::io::ErrorKind::TimedOut,
"daemon initial request frame read timed out",
)
}
#[cfg(unix)]
async fn read_initial_frame<R>(
stream: &mut R,
deadline: tokio::time::Instant,
) -> std::io::Result<Vec<u8>>
where
R: tokio::io::AsyncRead + Unpin,
{
if tokio::time::Instant::now() >= deadline {
return Err(initial_frame_timeout_error());
}
let raw = tokio::time::timeout_at(deadline, read_frame(stream))
.await
.map_err(|_| initial_frame_timeout_error())??;
if tokio::time::Instant::now() >= deadline {
return Err(initial_frame_timeout_error());
}
Ok(raw)
}
#[cfg(unix)]
pub async fn write_frame<W>(stream: &mut W, payload: &[u8]) -> std::io::Result<()>
where
W: tokio::io::AsyncWrite + Unpin,
{
if payload.len() > MAX_FRAME_BYTES {
return Err(std::io::Error::new(
std::io::ErrorKind::InvalidData,
format!(
"daemon frame of {} bytes exceeds {MAX_FRAME_BYTES} cap",
payload.len()
),
));
}
let len = (payload.len() as u32).to_be_bytes();
stream.write_all(&len).await?;
stream.write_all(payload).await?;
stream.flush().await?;
Ok(())
}
#[cfg(unix)]
#[async_trait]
pub trait DaemonDispatch: Clone + Send + Sync + 'static {
fn idle_retirement_blockers(&self) -> Vec<String> {
vec!["dispatcher_resource_inventory_unknown".to_owned()]
}
fn plan(&self, ops: &str) -> String;
#[allow(clippy::too_many_arguments)]
async fn dispatch(
&self,
ops: String,
presentation: Option<String>,
presentation_per_op: Option<Vec<Option<String>>>,
format: Option<String>,
format_per_op: Option<Vec<Option<String>>>,
from_wire: bool,
identity: Option<RequestIdentity>,
) -> Result<String, String>;
fn request_read_timeout(&self, _ops: &str) -> std::time::Duration {
khive_storage::request_read_timeout_from_env()
}
#[allow(clippy::too_many_arguments)]
async fn dispatch_with_error_detail(
&self,
ops: String,
presentation: Option<String>,
presentation_per_op: Option<Vec<Option<String>>>,
format: Option<String>,
format_per_op: Option<Vec<Option<String>>>,
from_wire: bool,
identity: Option<RequestIdentity>,
) -> Result<String, DaemonDispatchError> {
self.dispatch(
ops,
presentation,
presentation_per_op,
format,
format_per_op,
from_wire,
identity,
)
.await
.map_err(|message| DaemonDispatchError::new(message, None))
}
async fn warm_all(&self);
fn namespace(&self) -> &str;
fn config_id(&self) -> &str;
fn pool_for_checkpoint(&self) -> Option<Arc<ConnectionPool>> {
None
}
fn secondary_pools_for_checkpoint(&self) -> Vec<Arc<ConnectionPool>> {
Vec::new()
}
fn event_store_for_checkpoint(&self) -> Option<Arc<dyn khive_storage::EventStore>> {
None
}
}
#[cfg(unix)]
struct CheckpointTaskSpec {
pool: Arc<ConnectionPool>,
lifecycle_owner: Option<CheckpointLifecycleOwner>,
is_main: bool,
}
#[cfg(unix)]
fn checkpoint_task_specs(
main_pool: Option<Arc<ConnectionPool>>,
secondary_pools: Vec<Arc<ConnectionPool>>,
event_store: Option<Arc<dyn khive_storage::EventStore>>,
namespace: String,
) -> Vec<CheckpointTaskSpec> {
let mut tasks = Vec::with_capacity(usize::from(main_pool.is_some()) + secondary_pools.len());
if let Some(pool) = main_pool {
tasks.push(CheckpointTaskSpec {
pool,
lifecycle_owner: None,
is_main: true,
});
}
tasks.extend(secondary_pools.into_iter().map(|pool| CheckpointTaskSpec {
pool,
lifecycle_owner: None,
is_main: false,
}));
if let (Some(task), Some(event_store)) = (tasks.first_mut(), event_store) {
task.lifecycle_owner = Some(CheckpointLifecycleOwner::new(event_store, namespace));
}
tasks
}
static WARM_INDEX_HOST: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
pub fn mark_warm_index_host() {
WARM_INDEX_HOST.store(true, std::sync::atomic::Ordering::Release);
}
pub fn is_warm_index_host() -> bool {
WARM_INDEX_HOST.load(std::sync::atomic::Ordering::Acquire)
}
static BACKGROUND_TASKS: std::sync::OnceLock<Arc<std::sync::atomic::AtomicUsize>> =
std::sync::OnceLock::new();
fn background_tasks() -> &'static Arc<std::sync::atomic::AtomicUsize> {
BACKGROUND_TASKS.get_or_init(|| Arc::new(std::sync::atomic::AtomicUsize::new(0)))
}
static BACKGROUND_TASK_NAMES: std::sync::OnceLock<
std::sync::Mutex<std::collections::HashMap<&'static str, usize>>,
> = std::sync::OnceLock::new();
fn background_task_names_registry(
) -> &'static std::sync::Mutex<std::collections::HashMap<&'static str, usize>> {
BACKGROUND_TASK_NAMES.get_or_init(|| std::sync::Mutex::new(std::collections::HashMap::new()))
}
pub const UNNAMED_BACKGROUND_TASK: &str = "unnamed";
fn register_background_task_name(name: &'static str) {
let mut names = background_task_names_registry()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
*names.entry(name).or_insert(0) += 1;
}
fn release_background_task_name(name: &'static str) {
let mut names = background_task_names_registry()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(count) = names.get_mut(name) {
*count -= 1;
if *count == 0 {
names.remove(name);
}
}
}
pub fn background_task_names() -> Vec<String> {
let names = background_task_names_registry()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let mut out: Vec<String> = names.keys().map(|name| (*name).to_string()).collect();
out.sort();
out
}
#[cfg(unix)]
fn idle_retirement_blockers<D: DaemonDispatch>(dispatcher: &D) -> Vec<String> {
let mut blockers = dispatcher.idle_retirement_blockers();
if !khive_storage::tx_registry::snapshot().is_empty() {
blockers.push("open_sql_transaction".to_owned());
}
blockers.extend(
active_phase_names()
.into_iter()
.map(|name| format!("active_phase:{name}")),
);
let count = background_task_count();
let names = background_task_names_registry()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if names.values().sum::<usize>() != count {
blockers.push("tracked_worker_inventory_unsettled".to_owned());
}
for name in names.keys() {
if !matches!(
*name,
"wal_checkpoint" | "memory_ann_rotation_watch" | "knowledge_ann_rotation_watch"
) {
blockers.push(format!("unsettled_worker:{name}"));
}
}
blockers.sort();
blockers.dedup();
blockers
}
#[cfg(unix)]
async fn wait_for_idle<D: DaemonDispatch>(dispatcher: &D, lifecycle: &DaemonLifecycle) {
if lifecycle.options.lifetime == DaemonLifetime::Persistent {
std::future::pending::<()>().await;
}
loop {
if lifecycle.try_idle(|| idle_retirement_blockers(dispatcher)) {
return;
}
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
}
}
struct BackgroundTaskGuard {
counter: Arc<std::sync::atomic::AtomicUsize>,
name: &'static str,
}
impl Drop for BackgroundTaskGuard {
fn drop(&mut self) {
release_background_task_name(self.name);
self.counter
.fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
}
}
pub fn spawn_tracked_task<F, T>(fut: F) -> tokio::task::JoinHandle<T>
where
F: std::future::Future<Output = T> + Send + 'static,
T: Send + 'static,
{
spawn_named_tracked_task(UNNAMED_BACKGROUND_TASK, fut)
}
pub fn spawn_named_tracked_task<F, T>(name: &'static str, fut: F) -> tokio::task::JoinHandle<T>
where
F: std::future::Future<Output = T> + Send + 'static,
T: Send + 'static,
{
background_tasks().fetch_add(1, std::sync::atomic::Ordering::SeqCst);
register_background_task_name(name);
let guard = BackgroundTaskGuard {
counter: background_tasks().clone(),
name,
};
tokio::spawn(async move {
let _guard = guard;
fut.await
})
}
pub fn track_background_task<F>(fut: F)
where
F: std::future::Future<Output = ()> + Send + 'static,
{
track_named_background_task(UNNAMED_BACKGROUND_TASK, fut);
}
pub fn track_named_background_task<F>(name: &'static str, fut: F)
where
F: std::future::Future<Output = ()> + Send + 'static,
{
drop(spawn_named_tracked_task(name, fut));
}
pub fn background_task_count() -> usize {
background_tasks().load(std::sync::atomic::Ordering::SeqCst)
}
pub fn daemon_shutdown_token() -> tokio_util::sync::CancellationToken {
static TOKEN: std::sync::OnceLock<tokio_util::sync::CancellationToken> =
std::sync::OnceLock::new();
TOKEN
.get_or_init(tokio_util::sync::CancellationToken::new)
.clone()
}
static ACTIVE_PHASES: std::sync::OnceLock<
std::sync::Mutex<std::collections::HashMap<String, usize>>,
> = std::sync::OnceLock::new();
fn active_phases() -> &'static std::sync::Mutex<std::collections::HashMap<String, usize>> {
ACTIVE_PHASES.get_or_init(|| std::sync::Mutex::new(std::collections::HashMap::new()))
}
pub struct PhaseGuard {
name: String,
}
impl Drop for PhaseGuard {
fn drop(&mut self) {
let mut map = active_phases()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(count) = map.get_mut(&self.name) {
*count -= 1;
if *count == 0 {
map.remove(&self.name);
}
}
}
}
pub fn register_active_phase(name: &str) -> PhaseGuard {
let mut map = active_phases()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
*map.entry(name.to_string()).or_insert(0) += 1;
PhaseGuard {
name: name.to_string(),
}
}
pub fn active_phase_names() -> Vec<String> {
let map = active_phases()
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let mut names: Vec<String> = map.keys().cloned().collect();
names.sort();
names
}
#[cfg(unix)]
fn build_metrics_snapshot<D: DaemonDispatch>(dispatcher: &D) -> MetricsSnapshot {
let open_tx_count = khive_storage::tx_registry::snapshot().len();
let (oldest_pinned_tx_micros, oldest_pinned_tx_label) =
match khive_storage::tx_registry::oldest() {
Some((_id, age, label)) => (Some(age.as_micros() as u64), label),
None => (None, None),
};
let checkpoint_pool = dispatcher.pool_for_checkpoint();
let mut secondary_index = 0;
let wal_checkpoint_stores = checkpoint_task_specs(
checkpoint_pool.clone(),
dispatcher.secondary_pools_for_checkpoint(),
None,
String::new(),
)
.into_iter()
.map(|task| {
let (store_id, role) = if task.is_main {
("main".to_string(), "main".to_string())
} else {
let store_id = format!("secondary:{secondary_index}");
secondary_index += 1;
(store_id, "secondary".to_string())
};
CheckpointStoreMetrics {
store_id,
role,
database: task
.pool
.canonical_path()
.and_then(std::path::Path::file_name)
.map(|name| name.to_string_lossy().into_owned()),
timing: khive_db::checkpoint::checkpoint_timing(&task.pool),
}
})
.collect();
let routine_wal = checkpoint_pool
.as_deref()
.and_then(khive_db::checkpoint::routine_wal_observation);
let writer_stages = checkpoint_pool
.as_deref()
.and_then(khive_db::writer_task::last_writer_stage_observation);
let (write_queue_depth, write_queue_capacity) = checkpoint_pool
.as_ref()
.and_then(|pool| pool.writer_task_handle().ok().flatten())
.map(|handle| (Some(handle.queue_depth()), Some(handle.capacity())))
.unwrap_or((None, None));
MetricsSnapshot {
lifecycle: None,
wal_pages: routine_wal.as_ref().map(|sample| sample.log_frames),
wal_log_frames: routine_wal.as_ref().map(|sample| sample.log_frames),
wal_checkpointed_frames: routine_wal
.as_ref()
.map(|sample| sample.checkpointed_frames),
wal_pending_frames: routine_wal.as_ref().map(|sample| sample.pending_frames),
wal_physical_bytes: routine_wal
.as_ref()
.and_then(|sample| sample.physical_wal_bytes),
wal_observed_at_unix_ms: routine_wal
.as_ref()
.map(|sample| sample.observed_at_unix_ms),
wal_checkpoint_stores,
wal_truncate_attempts: khive_db::checkpoint::truncate_attempts(),
wal_truncate_consecutive_failures: khive_db::checkpoint::truncate_consecutive_failures(),
wal_checkpoint_skipped_ticks: khive_db::checkpoint::checkpoint_skipped_ticks(),
wal_checkpoint_consecutive_skips: khive_db::checkpoint::checkpoint_consecutive_skips(),
wal_checkpoint_last_skip_wal_pages: khive_db::checkpoint::checkpoint_last_skip_wal_pages(),
oldest_pinned_tx_micros,
oldest_pinned_tx_label,
open_tx_count,
write_queue_depth,
write_queue_capacity,
write_last_queue_wait_micros: writer_stages
.as_ref()
.map(|sample| sample.queue_wait_micros),
write_last_transaction_acquire_micros: writer_stages
.as_ref()
.map(|sample| sample.transaction_acquire_micros),
write_last_body_micros: writer_stages.as_ref().map(|sample| sample.body_micros),
write_last_commit_micros: writer_stages.as_ref().map(|sample| sample.commit_micros),
write_last_total_micros: writer_stages.as_ref().map(|sample| sample.total_micros),
write_last_observed_at_unix_ms: writer_stages
.as_ref()
.map(|sample| sample.observed_at_unix_ms),
connections: None,
recall_ledger: Some(recall_ledger_snapshot()),
}
}
#[cfg(unix)]
async fn write_response_frame<W>(stream: &mut W, payload: &[u8]) -> std::io::Result<()>
where
W: tokio::io::AsyncWrite + Unpin,
{
tokio::time::timeout(INITIAL_FRAME_READ_TIMEOUT, write_frame(stream, payload))
.await
.map_err(|_| {
std::io::Error::new(
std::io::ErrorKind::TimedOut,
"daemon response write timed out",
)
})?
}
#[cfg(unix)]
async fn wait_for_peer_disconnect(read: &mut tokio::net::unix::OwnedReadHalf) {
let mut byte = [0u8; 1];
let _ = read.read(&mut byte).await;
}
#[cfg(all(unix, test))]
async fn handle_conn<D: DaemonDispatch>(stream: UnixStream, dispatcher: D) {
handle_conn_with_shutdown(
stream,
dispatcher,
None,
tokio::time::Instant::now() + INITIAL_FRAME_READ_TIMEOUT,
)
.await;
}
#[cfg(all(unix, feature = "fault-injection"))]
#[doc(hidden)]
pub async fn handle_conn_for_test<D: DaemonDispatch>(stream: UnixStream, dispatcher: D) {
handle_conn_with_shutdown(
stream,
dispatcher,
None,
tokio::time::Instant::now() + INITIAL_FRAME_READ_TIMEOUT,
)
.await;
}
#[cfg(unix)]
fn plan_frame_companion(raw: &[u8]) -> Option<&'static str> {
let value: serde_json::Value = serde_json::from_slice(raw).ok()?;
if value.get("plan").and_then(serde_json::Value::as_bool) != Some(true)
|| value
.get("protocol_version")
.and_then(serde_json::Value::as_u64)
!= Some(u64::from(PROTOCOL_VERSION))
{
return None;
}
[
"presentation",
"presentation_per_op",
"format",
"format_per_op",
"request_id",
]
.into_iter()
.find(|field| value.get(*field).is_some())
}
#[cfg(all(
unix,
any(test, feature = "fault-injection", feature = "test-internals")
))]
async fn handle_conn_with_shutdown<D: DaemonDispatch>(
stream: UnixStream,
dispatcher: D,
shutdown: Option<tokio::sync::watch::Receiver<bool>>,
initial_frame_deadline: tokio::time::Instant,
) {
handle_conn_with_lifecycle(stream, dispatcher, shutdown, initial_frame_deadline, None).await;
}
#[cfg(unix)]
async fn handle_conn_with_lifecycle<D: DaemonDispatch>(
mut stream: UnixStream,
dispatcher: D,
shutdown: Option<tokio::sync::watch::Receiver<bool>>,
initial_frame_deadline: tokio::time::Instant,
lifecycle: Option<Arc<DaemonLifecycle>>,
) {
let mut ordinary_admission = None;
let production_shutdown = shutdown.is_some();
let handover_peer_allowed = peer_uid(&stream)
.ok()
.is_some_and(|uid| uid == unsafe { libc::geteuid() } as u32);
let (local_shutdown_tx, local_shutdown_rx) = tokio::sync::watch::channel(false);
let shutdown = shutdown.unwrap_or(local_shutdown_rx);
let _local_shutdown_tx = local_shutdown_tx;
let raw = match read_initial_frame(&mut stream, initial_frame_deadline).await {
Ok(r) => r,
Err(e) => {
tracing::debug!(error = %e, "failed to read daemon request frame");
return;
}
};
#[derive(Deserialize)]
struct SupervisorRequestEnvelope {
#[serde(flatten)]
frame: DaemonRequestFrame,
#[serde(default)]
supervisor_handover: bool,
}
let decoded: Result<SupervisorRequestEnvelope, _> = serde_json::from_slice(&raw);
if decoded.as_ref().ok().is_none_or(|item| item.frame.plan) {
if let Some(field) = plan_frame_companion(&raw) {
let response = DaemonResponseFrame {
ok: false,
result: None,
error: Some(format!(
"invalid_params: plan=true cannot be combined with {field}"
)),
error_detail: Some(serde_json::json!({
"kind": "protocol",
"code": "invalid_params",
"message": format!("plan=true cannot be combined with {field}"),
"domain_disposition": crate::DomainDisposition::NotCommitted.as_str(),
})),
namespace_mismatch: false,
config_mismatch: false,
served_config_id: Some(dispatcher.config_id().to_string()),
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: None,
};
if let Ok(payload) = serde_json::to_vec(&response) {
if let Err(error) = write_response_frame(&mut stream, &payload).await {
tracing::debug!(%error, "failed to write plan envelope refusal");
}
}
return;
}
}
let (frame, handover_requested) = match decoded {
Ok(item) => (item.frame, item.supervisor_handover),
Err(e) => {
tracing::debug!(error = %e, "failed to decode daemon request frame");
return;
}
};
let supervisor_probe = frame.probe_only;
let handover_accepted = handover_requested
&& supervisor_probe
&& production_shutdown
&& handover_peer_allowed
&& frame.protocol_version == PROTOCOL_VERSION
&& !frame.plan
&& !frame.metrics_only
&& frame.ops.is_empty()
&& read_supervisor_marker_claim().is_some_and(|(pid, _)| pid != std::process::id());
let (mut peer_read, mut peer_write) = stream.into_split();
let served_config_id = Some(dispatcher.config_id().to_string());
let resp = if frame.protocol_version != PROTOCOL_VERSION {
let msg = format!(
"daemon protocol mismatch: client={} daemon={} — \
rebuild/update the client binary (make local)",
frame.protocol_version, PROTOCOL_VERSION,
);
tracing::warn!(
client_version = frame.protocol_version,
daemon_version = PROTOCOL_VERSION,
"daemon protocol version mismatch"
);
DaemonResponseFrame {
ok: false,
result: None,
error: Some(msg.clone()),
error_detail: Some(serde_json::json!({
"kind": "protocol",
"code": "version_mismatch",
"message": msg,
"domain_disposition": crate::DomainDisposition::Unknown.as_str(),
})),
namespace_mismatch: false,
config_mismatch: false,
served_config_id,
version_mismatch: frame.protocol_version > PROTOCOL_VERSION,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: frame.request_id,
}
} else if handover_requested {
if handover_accepted {
DaemonResponseFrame {
ok: true,
result: None,
error: None,
error_detail: None,
namespace_mismatch: false,
config_mismatch: false,
served_config_id,
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: frame.request_id,
}
} else {
DaemonResponseFrame {
ok: false,
result: None,
error: Some("supervisor handover refused".to_string()),
error_detail: Some(serde_json::json!({
"kind": "protocol",
"code": "supervisor_handover_refused",
"message": "supervisor handover refused",
"domain_disposition": crate::DomainDisposition::NotCommitted.as_str(),
})),
namespace_mismatch: false,
config_mismatch: false,
served_config_id,
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: frame.request_id,
}
}
} else if frame.metrics_only && !frame.plan {
DaemonResponseFrame {
ok: true,
result: None,
error: None,
error_detail: None,
namespace_mismatch: false,
config_mismatch: false,
served_config_id,
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: Some({
let mut metrics = build_metrics_snapshot(&dispatcher);
metrics.lifecycle = lifecycle.as_ref().map(|state| {
let mut snapshot = state.snapshot();
snapshot.idle_blockers = idle_retirement_blockers(&dispatcher);
snapshot
});
metrics.connections = lifecycle.as_ref().map(|state| state.connections.snapshot());
metrics
}),
request_id: frame.request_id,
}
} else if !config_ids_compatible(&frame.config_id, dispatcher.config_id()) {
DaemonResponseFrame {
ok: false,
result: None,
error: None,
error_detail: Some(serde_json::json!({
"kind": "protocol",
"code": "config_mismatch",
"message": "daemon configuration does not match the request",
"domain_disposition": crate::DomainDisposition::NotCommitted.as_str(),
})),
namespace_mismatch: false,
config_mismatch: true,
served_config_id,
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: frame.request_id,
}
} else if frame.plan {
DaemonResponseFrame {
ok: true,
result: Some(dispatcher.plan(&frame.ops)),
error: None,
error_detail: None,
namespace_mismatch: false,
config_mismatch: false,
served_config_id,
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: None,
}
} else if frame.probe_only {
DaemonResponseFrame {
ok: true,
result: None,
error: None,
error_detail: None,
namespace_mismatch: false,
config_mismatch: false,
served_config_id,
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: frame.request_id,
}
} else {
if let Some(lifecycle) = &lifecycle {
ordinary_admission = lifecycle.admit();
if ordinary_admission.is_none() {
let refusal = DaemonResponseFrame {
ok: false,
result: None,
error: Some("daemon is draining; request was not admitted".to_owned()),
error_detail: Some(serde_json::json!({
"kind": "runtime", "code": "daemon_draining",
"domain_disposition": crate::DomainDisposition::NotCommitted.as_str(),
})),
request_id: frame.request_id,
daemon_protocol_version: PROTOCOL_VERSION,
namespace_mismatch: false,
config_mismatch: false,
served_config_id: Some(dispatcher.config_id().to_owned()),
version_mismatch: false,
metrics: None,
};
if let Ok(payload) = serde_json::to_vec(&refusal) {
let _ = write_response_frame(&mut peer_write, &payload).await;
}
return;
}
}
let identity = RequestIdentity {
namespace: frame.namespace.clone(),
actor_id: frame.actor_id.clone(),
visible_namespaces: frame.visible_namespaces.clone(),
process_ref: frame.process_ref.clone(),
request_id: frame.request_id,
};
tracing::debug!(
request_id = frame.request_id,
"daemon RequestIdentity constructed"
);
let (read_cancel_tx, read_cancel_rx) = tokio::sync::watch::channel(false);
let read_timeout = dispatcher.request_read_timeout(&frame.ops);
let excluded_embedder_names =
config_id_extra_embedder_exclusions(&frame.config_id, dispatcher.config_id())
.expect("compatible configuration ids must expose their extra embedder sets");
let dispatch = crate::runtime::scope_request_embedder_exclusions(
excluded_embedder_names,
khive_storage::scope_request_read_cancellation(
shutdown,
khive_storage::scope_request_read_cancellation(
read_cancel_rx,
khive_storage::scope_request_read_deadline(
read_timeout,
dispatcher.dispatch_with_error_detail(
frame.ops,
frame.presentation,
frame.presentation_per_op,
frame.format,
frame.format_per_op,
frame.from_wire,
Some(identity),
),
),
),
),
);
tokio::pin!(dispatch);
let dispatch_result = tokio::select! {
result = &mut dispatch => result,
_ = wait_for_peer_disconnect(&mut peer_read) => {
let _ = read_cancel_tx.send(true);
dispatch.await
}
};
match dispatch_result {
Ok(result) => {
let (result, detail) = take_daemon_lexical_timeout_marker(result);
DaemonResponseFrame {
ok: true,
result: Some(result),
error: None,
error_detail: detail,
namespace_mismatch: false,
config_mismatch: false,
served_config_id,
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: frame.request_id,
}
}
Err(error) => {
let error = DaemonDispatchError::new(error.message, Some(error.error_detail));
DaemonResponseFrame {
ok: false,
result: None,
error: Some(error.message),
error_detail: Some(error.error_detail),
namespace_mismatch: false,
config_mismatch: false,
served_config_id,
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: frame.request_id,
}
}
}
};
let payload = if supervisor_probe {
serde_json::to_value(&resp).and_then(|mut value| {
if let Some(claim) = current_supervisor_claim() {
value["supervisor_claim"] = serde_json::Value::String(claim);
}
if handover_accepted {
value["supervisor_handover_accepted"] = serde_json::Value::Bool(true);
}
serde_json::to_vec(&value)
})
} else {
serde_json::to_vec(&resp)
};
let mut handover_ack_written = false;
match payload {
Ok(payload) => {
if payload.len() > MAX_FRAME_BYTES {
tracing::warn!(
bytes = payload.len(),
limit = MAX_FRAME_BYTES,
"daemon response exceeds MAX_FRAME_BYTES; sending explicit error frame"
);
let message = format!(
"response too large: {} bytes exceeds {} byte IPC cap",
payload.len(),
MAX_FRAME_BYTES,
);
let err_resp = DaemonResponseFrame {
ok: false,
result: None,
error: Some(message.clone()),
error_detail: Some(serde_json::json!({
"kind": "transport",
"code": "response_frame_size_limit",
"message": message,
"domain_disposition": crate::DomainDisposition::Unknown.as_str(),
})),
namespace_mismatch: false,
config_mismatch: false,
served_config_id: resp.served_config_id,
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: resp.request_id,
};
if let Ok(err_payload) = serde_json::to_vec(&err_resp) {
if let Err(e) = write_response_frame(&mut peer_write, &err_payload).await {
tracing::debug!(error = %e, "failed to write oversized-response error frame");
}
}
} else {
match write_response_frame(&mut peer_write, &payload).await {
Ok(()) => handover_ack_written = true,
Err(e) => tracing::debug!(error = %e, "failed to write daemon response frame"),
}
}
}
Err(e) => tracing::warn!(error = %e, "failed to serialize daemon response frame"),
}
if handover_accepted && handover_ack_written {
if unsafe { libc::raise(libc::SIGTERM) } != 0 {
tracing::error!(error = %std::io::Error::last_os_error(), "self-directed handover signal failed");
}
}
drop(ordinary_admission);
}
#[cfg(unix)]
struct ActiveConnectionGuard {
active: Arc<std::sync::atomic::AtomicUsize>,
}
#[cfg(unix)]
impl ActiveConnectionGuard {
fn claim(active: Arc<std::sync::atomic::AtomicUsize>) -> Self {
active.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Self { active }
}
}
#[cfg(unix)]
impl Drop for ActiveConnectionGuard {
fn drop(&mut self) {
self.active
.fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
}
}
#[cfg(unix)]
fn spawn_connection_task<F>(
active: Arc<std::sync::atomic::AtomicUsize>,
future: F,
) -> tokio::task::JoinHandle<()>
where
F: std::future::Future<Output = ()> + Send + 'static,
{
let guard = ActiveConnectionGuard::claim(active);
tokio::spawn(async move {
let _guard = guard;
future.await;
})
}
#[cfg(unix)]
pub async fn run_daemon<D: DaemonDispatch>(dispatcher: D) -> anyhow::Result<()> {
let boot_guard = Some(acquire_daemon_boot_guard()?);
run_daemon_with_boot_guard_inner(
dispatcher,
boot_guard,
false,
DaemonOptions::default(),
|_| DaemonStartupReport::default(),
)
.await
}
#[cfg(all(unix, any(test, feature = "fault-injection")))]
#[doc(hidden)]
pub async fn run_daemon_in_process_test<D: DaemonDispatch>(dispatcher: D) -> anyhow::Result<()> {
let boot_guard = Some(acquire_daemon_boot_guard()?);
run_daemon_with_boot_guard_inner(
dispatcher,
boot_guard,
true,
DaemonOptions::default(),
|_| DaemonStartupReport::default(),
)
.await
}
#[cfg(unix)]
#[derive(Clone, Copy, PartialEq, Eq)]
enum RendezvousPathRole {
Socket,
PidFile,
}
#[cfg(unix)]
impl RendezvousPathRole {
fn env_name(self) -> &'static str {
match self {
Self::Socket => SOCKET_PATH_ENV,
Self::PidFile => PID_PATH_ENV,
}
}
fn directory_name(self) -> &'static str {
match self {
Self::Socket => "socket directory",
Self::PidFile => "PID-file directory",
}
}
fn path_name(self) -> &'static str {
match self {
Self::Socket => "socket path",
Self::PidFile => "PID-file path",
}
}
fn path_component_name(self) -> &'static str {
match self {
Self::Socket => "socket-path",
Self::PidFile => "PID-file-path",
}
}
}
#[cfg(unix)]
pub(crate) fn ensure_socket_dir_is_trusted(parent: &std::path::Path) -> anyhow::Result<()> {
let daemon_euid = unsafe { libc::geteuid() } as u32;
ensure_rendezvous_dir_is_trusted(parent, RendezvousPathRole::Socket, daemon_euid, true)
}
#[cfg(unix)]
pub fn ensure_pid_file_dir_is_trusted(pid_file: &std::path::Path) -> anyhow::Result<()> {
let parent = pid_file
.parent()
.filter(|parent| !parent.as_os_str().is_empty())
.unwrap_or_else(|| std::path::Path::new("."));
let daemon_euid = unsafe { libc::geteuid() } as u32;
ensure_rendezvous_dir_is_trusted(parent, RendezvousPathRole::PidFile, daemon_euid, false)
}
#[cfg(unix)]
fn ensure_rendezvous_dir_is_trusted(
parent: &std::path::Path,
role: RendezvousPathRole,
daemon_euid: u32,
repair_owned_default: bool,
) -> anyhow::Result<()> {
let env_name = role.env_name();
let directory_name = role.directory_name();
if repair_owned_default && parent == khive_dir() {
std::fs::set_permissions(parent, std::fs::Permissions::from_mode(0o700)).map_err(|e| {
anyhow::anyhow!(
"refusing to start: cannot chmod 0700 {}: {e}. The khive directory must be \
owner-only as the {directory_name} for {env_name}; it is part of the \
same-uid guarantee this daemon enforces.",
parent.display()
)
})?;
return ensure_rendezvous_path_is_swap_resistant(parent, daemon_euid, role);
}
let meta = std::fs::metadata(parent).map_err(|e| {
anyhow::anyhow!(
"refusing to start: cannot stat {directory_name} {} for {env_name}: {e}. \
It gates rendezvous-path safety, and unreadable metadata is not a passing state.",
parent.display()
)
})?;
use std::os::unix::fs::MetadataExt;
let owner = meta.uid();
if owner != daemon_euid && owner != 0 {
anyhow::bail!(
"refusing to start: {directory_name} {} for {env_name} is owned by uid {owner}, \
not this daemon's uid ({daemon_euid}) or root. A directory owner can replace \
the rendezvous path regardless of mode bits. Point {env_name} at a directory \
you own, or unset it for the default.",
parent.display()
);
}
let mode = meta.permissions().mode();
if mode & 0o022 != 0 {
anyhow::bail!(
"refusing to start: {directory_name} {} for {env_name} is mode {:04o} — writable \
by group or other, so another local user could replace the rendezvous path. \
Use a directory only you can write, or unset {env_name} for the default. \
This daemon is not changing the permissions of a directory it does not own.",
parent.display(),
mode & 0o7777
);
}
ensure_rendezvous_path_is_swap_resistant(parent, daemon_euid, role)
}
#[cfg(all(unix, test))]
fn ensure_socket_path_is_swap_resistant(
parent: &std::path::Path,
daemon_euid: u32,
) -> anyhow::Result<()> {
ensure_rendezvous_path_is_swap_resistant(parent, daemon_euid, RendezvousPathRole::Socket)
}
#[cfg(unix)]
fn ensure_rendezvous_path_is_swap_resistant(
parent: &std::path::Path,
daemon_euid: u32,
role: RendezvousPathRole,
) -> anyhow::Result<()> {
use std::os::unix::fs::MetadataExt;
let env_name = role.env_name();
let directory_name = role.directory_name();
let path_name = role.path_name();
let component_name = role.path_component_name();
let absolute = if parent.is_absolute() {
parent.to_path_buf()
} else {
std::env::current_dir()
.map_err(|e| {
anyhow::anyhow!(
"refusing to start: cannot resolve the working directory to absolutize \
{directory_name} {} for {env_name}: {e}.",
parent.display()
)
})?
.join(parent)
};
fn push_components(stack: &mut Vec<std::ffi::OsString>, path: &std::path::Path) {
let components: Vec<_> = path
.components()
.map(|c| c.as_os_str().to_os_string())
.collect();
stack.extend(components.into_iter().rev());
}
let mut stack: Vec<std::ffi::OsString> = Vec::new();
push_components(&mut stack, &absolute);
let mut resolved = std::path::PathBuf::new();
let mut symlinks_followed = 0u32;
while let Some(component) = stack.pop() {
if component == "/" {
resolved = std::path::PathBuf::from("/");
continue;
}
if component == "." {
continue;
}
if component == ".." {
resolved.pop();
continue;
}
let candidate = resolved.join(&component);
let meta = std::fs::symlink_metadata(&candidate).map_err(|e| {
anyhow::anyhow!(
"refusing to start: cannot stat {component_name} component {} for {env_name}: \
{e}. An unreadable component is not a passing one.",
candidate.display()
)
})?;
let owner = meta.uid();
if meta.file_type().is_symlink() {
symlinks_followed += 1;
if symlinks_followed > 40 {
anyhow::bail!(
"refusing to start: {path_name} for {env_name} resolves through more than \
40 symlinks at {} — treating this as a loop.",
candidate.display()
);
}
if owner != daemon_euid && owner != 0 {
anyhow::bail!(
"refusing to start: {component_name} symlink component {} for {env_name} \
is owned by uid {owner}, not this daemon's uid ({daemon_euid}) or root — \
its owner could retarget it after this check and re-root the {path_name}. \
Point {env_name} somewhere trusted end to end, or unset it for the default.",
candidate.display()
);
}
let target = std::fs::read_link(&candidate).map_err(|e| {
anyhow::anyhow!(
"refusing to start: cannot read {component_name} symlink component {} \
for {env_name}: {e}.",
candidate.display()
)
})?;
push_components(&mut stack, &target);
continue;
}
if meta.is_dir() {
let mode = meta.permissions().mode();
let sticky = mode & 0o1000 != 0;
if owner != daemon_euid && owner != 0 {
anyhow::bail!(
"refusing to start: {component_name} ancestor {} for {env_name} is owned by \
uid {owner}, not this daemon's uid ({daemon_euid}) or root — its owner \
could rename the next path component and re-root the {path_name}. Point \
{env_name} somewhere trusted end to end, or unset it for the default.",
candidate.display()
);
}
if mode & 0o022 != 0 && !sticky {
anyhow::bail!(
"refusing to start: {component_name} ancestor {} for {env_name} is mode \
{:04o} — writable by group or other without the sticky bit, so another \
local user could rename the next path component and re-root the {path_name}. \
Point {env_name} somewhere trusted end to end, or unset it for the default.",
candidate.display(),
mode & 0o7777
);
}
resolved = candidate;
continue;
}
anyhow::bail!(
"refusing to start: {component_name} component {} for {env_name} is neither a \
directory nor a symlink — the {path_name} cannot traverse it.",
candidate.display()
);
}
Ok(())
}
#[cfg(unix)]
pub async fn run_daemon_with_boot_guard<D: DaemonDispatch>(
dispatcher: D,
boot_guard: Option<std::fs::File>,
) -> anyhow::Result<()> {
run_daemon_with_boot_guard_inner(
dispatcher,
boot_guard,
false,
DaemonOptions::default(),
|_| DaemonStartupReport::default(),
)
.await
}
#[cfg(unix)]
pub async fn run_daemon_with_boot_guard_and_start<D, F>(
dispatcher: D,
boot_guard: Option<std::fs::File>,
start: F,
) -> anyhow::Result<()>
where
D: DaemonDispatch,
F: FnOnce(&D) + Send,
{
run_daemon_with_options_and_boot_guard_and_start(
dispatcher,
boot_guard,
DaemonOptions::default(),
|dispatcher| {
start(dispatcher);
DaemonStartupReport::default()
},
)
.await
}
#[cfg(unix)]
pub async fn run_daemon_with_options_and_boot_guard_and_start<D, F>(
dispatcher: D,
boot_guard: Option<std::fs::File>,
options: DaemonOptions,
start: F,
) -> anyhow::Result<()>
where
D: DaemonDispatch,
F: FnOnce(&D) -> DaemonStartupReport + Send,
{
anyhow::ensure!(
!options.idle_interval.is_zero(),
"daemon idle interval must be positive"
);
run_daemon_with_boot_guard_inner(dispatcher, boot_guard, false, options, start).await
}
#[cfg(unix)]
async fn run_daemon_with_boot_guard_inner<D, F>(
dispatcher: D,
boot_guard: Option<std::fs::File>,
allow_same_process_incumbent: bool,
options: DaemonOptions,
start: F,
) -> anyhow::Result<()>
where
D: DaemonDispatch,
F: FnOnce(&D) -> DaemonStartupReport + Send,
{
struct ComponentTeardown;
impl Drop for ComponentTeardown {
fn drop(&mut self) {
daemon_shutdown_token().cancel();
}
}
let _component_teardown = ComponentTeardown;
ensure_rendezvous_overrides_paired()?;
let sock = socket_path();
let pid_file = pid_path();
let socket_parent = sock.parent();
let pid_parent = pid_file.parent();
if let Some(parent) = socket_parent {
std::fs::create_dir_all(parent)?;
ensure_socket_dir_is_trusted(parent)?;
}
if pid_parent != socket_parent {
ensure_pid_file_dir_is_trusted(&pid_file)?;
}
let _startup_lock = boot_guard;
match cleanup_stale_daemon(
&sock,
&pid_file,
allow_same_process_incumbent,
dispatcher.config_id(),
)
.await
{
Incumbent::Serving(incumbent_pid) => {
tracing::error!(
pid = incumbent_pid,
socket = ?sock,
"refusing to start: a khived instance is already serving this socket"
);
anyhow::bail!(
"refusing to start: khived is already running as pid {incumbent_pid}, \
serving socket {}. Stop that instance first if you intend to replace it.",
sock.display()
);
}
Incumbent::Live(incumbent_pid) => {
tracing::error!(
pid = incumbent_pid,
socket = ?sock,
"refusing to start: a live process owns the PID file but no khived answered"
);
anyhow::bail!(
"refusing to start: pid {incumbent_pid} owns the daemon PID file and is alive, \
but nothing answered the khived protocol on {}. It may be draining. Nothing \
was removed; stop that process first if you intend to replace it.",
sock.display()
);
}
Incumbent::Stale => {}
}
let mut sigterm = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())?;
let mut sigint = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::interrupt())?;
let pid_file_guard = match write_pid_file_exclusive(&pid_file) {
Ok(guard) => guard,
Err(e) if e.kind() == std::io::ErrorKind::AlreadyExists => {
if pid_file_names_a_reachable_daemon(
&pid_file,
&sock,
allow_same_process_incumbent,
dispatcher.config_id(),
)
.await
{
tracing::info!(
"a replacement khived already claimed the pid/socket rendezvous; exiting"
);
return Ok(());
}
anyhow::bail!(
"failed to claim daemon pid file at {pid_file:?}: it already exists \
and does not name a reachable daemon"
);
}
Err(e) => return Err(e.into()),
};
let listener = match UnixListener::bind(&sock) {
Ok(listener) => listener,
Err(e) => {
remove_pid_file_if_owned(&pid_file, &pid_file_guard);
return Err(e.into());
}
};
if let Err(e) = std::fs::set_permissions(&sock, std::fs::Permissions::from_mode(0o600)) {
drop(listener);
let _ = std::fs::remove_file(&sock);
remove_pid_file_if_owned(&pid_file, &pid_file_guard);
return Err(anyhow::anyhow!(
"refusing to start: cannot chmod 0600 {}: {e}. The daemon socket must be owner-only \
— it is half of the single-principal guarantee this daemon enforces.",
sock.display()
));
}
let bound_identity = socket_identity(&sock);
let lifecycle = Arc::new(DaemonLifecycle::new(options, start(&dispatcher)));
drop(_startup_lock);
tracing::info!(
socket = ?sock,
pid = std::process::id(),
source_revision = crate::BUILD_INFO.source_revision,
build_time = crate::BUILD_INFO.build_time,
"khived listening"
);
{
let warm = dispatcher.clone();
track_named_background_task("daemon_warmup", async move {
warm.warm_all().await;
});
}
let (checkpoint_shutdown_tx, checkpoint_shutdown_rx) = tokio::sync::watch::channel(());
let checkpoint_tasks = checkpoint_task_specs(
dispatcher.pool_for_checkpoint(),
dispatcher.secondary_pools_for_checkpoint(),
dispatcher.event_store_for_checkpoint(),
dispatcher.namespace().to_string(),
);
if !checkpoint_tasks.is_empty() {
let cfg = CheckpointConfig::from_env();
let checkpoint_task_count = checkpoint_tasks.len();
for task in checkpoint_tasks {
track_named_background_task(
"wal_checkpoint",
run_checkpoint_task(
task.pool,
cfg.clone(),
task.lifecycle_owner,
checkpoint_shutdown_rx.clone(),
task.is_main,
),
);
}
tracing::info!(checkpoint_task_count, "WAL checkpoint task(s) started");
}
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let connection_tasks = Arc::new(std::sync::Mutex::new(
Vec::<tokio::task::JoinHandle<()>>::new(),
));
let (request_shutdown_tx, request_shutdown_rx) = tokio::sync::watch::channel(false);
let shutdown = async {
tokio::select! {
_ = sigterm.recv() => tracing::info!("received SIGTERM"),
_ = sigint.recv() => tracing::info!("received SIGINT"),
}
for signal in [libc::SIGTERM, libc::SIGINT] {
if unsafe { libc::signal(signal, libc::SIG_DFL) } == libc::SIG_ERR {
return Err(std::io::Error::last_os_error());
}
}
Ok::<(), std::io::Error>(())
};
tokio::pin!(shutdown);
lifecycle.ready();
let daemon_euid = unsafe { libc::geteuid() } as u32;
let reason = tokio::select! {
_ = async {
let mut accept_error_backoff = None;
let mut last_accept_error_log: Option<std::time::Instant> = None;
loop {
match listener.accept().await {
Ok((mut stream, _)) => {
let initial_frame_deadline =
tokio::time::Instant::now() + INITIAL_FRAME_READ_TIMEOUT;
accept_error_backoff = None;
last_accept_error_log = None;
match peer_uid(&stream) {
Ok(peer) if uid_is_permitted(peer, daemon_euid) => {}
Ok(peer) => {
tracing::error!(
peer_uid = peer,
daemon_euid,
"refusing connection from a foreign uid: this daemon accepts \
only peers running as its own uid"
);
drop(stream);
continue;
}
Err(e) => {
tracing::error!(
error = %e,
"refusing connection: cannot read peer credentials, so \
same-uid cannot be proven"
);
drop(stream);
continue;
}
}
let Some(permit) = admit_or_refuse_busy(
&lifecycle.connections,
&mut stream,
dispatcher.config_id(),
)
.await
else {
continue;
};
let d = dispatcher.clone();
let shutdown = request_shutdown_rx.clone();
let lifecycle = Arc::clone(&lifecycle);
let handle = spawn_connection_task(Arc::clone(&active), async move {
let _permit = permit;
handle_conn_with_lifecycle(
stream,
d,
Some(shutdown),
initial_frame_deadline,
Some(lifecycle),
)
.await;
});
let mut tasks = connection_tasks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
tasks.retain(|task| !task.is_finished());
tasks.push(handle);
}
Err(e) => {
let delay = next_accept_error_backoff(accept_error_backoff);
accept_error_backoff = Some(delay);
let capacity_exhausted = matches!(
e.raw_os_error(),
Some(libc::EMFILE) | Some(libc::ENFILE)
);
if last_accept_error_log.is_none_or(|last| {
last.elapsed() >= std::time::Duration::from_secs(30)
}) {
tracing::error!(
error = %e,
capacity_exhausted,
retry_ms = delay.as_millis(),
"daemon accept failed; retrying with bounded backoff"
);
last_accept_error_log = Some(std::time::Instant::now());
}
tokio::time::sleep(delay).await;
}
}
}
} => DaemonShutdownReason::Signal,
result = &mut shutdown => { result?; DaemonShutdownReason::Signal },
_ = wait_for_idle(&dispatcher, &lifecycle) => DaemonShutdownReason::Idle,
};
lifecycle.draining(reason);
drop(listener);
let _ = checkpoint_shutdown_tx.send(());
if reason == DaemonShutdownReason::Signal {
let _ = request_shutdown_tx.send(true);
}
daemon_shutdown_token().cancel();
let drained = if reason == DaemonShutdownReason::Idle {
tokio::select! {
_ = drain_for_idle(&active, drain_timeout()) => true,
result = &mut shutdown => {
result?;
lifecycle.draining(DaemonShutdownReason::Signal);
let _ = request_shutdown_tx.send(true);
drain(&active).await
}
}
} else {
drain(&active).await
};
let tasks = {
let mut retained = connection_tasks
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
std::mem::take(&mut *retained)
};
finish_connection_tasks(tasks, drained).await;
match acquire_recovery_lock() {
Some(_shutdown_lock) => {
shutdown_cleanup_if_owned(&sock, &pid_file, bound_identity);
}
None => {
tracing::warn!(
"could not acquire recovery lock for shutdown cleanup; \
skipping unlink to avoid deleting a replacement daemon's paths"
);
}
}
lifecycle.stopped();
tracing::info!("khived stopped");
Ok(())
}
#[cfg(unix)]
fn shutdown_cleanup_if_owned(
sock: &std::path::Path,
pid_file: &std::path::Path,
bound_identity: Option<SocketIdentity>,
) -> bool {
let pid_is_ours = std::fs::read_to_string(pid_file)
.ok()
.and_then(|s| s.trim().parse::<u32>().ok())
== Some(std::process::id());
let socket_is_ours = bound_identity.is_some() && socket_identity(sock) == bound_identity;
if pid_is_ours && socket_is_ours {
let _ = std::fs::remove_file(sock);
let _ = std::fs::remove_file(pid_file);
true
} else {
tracing::warn!(
socket = ?sock,
pid_file = ?pid_file,
"skipping shutdown cleanup — a replacement daemon already owns this socket/PID"
);
false
}
}
#[cfg(unix)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum PidLiveness {
Alive,
Dead,
PermissionDenied,
}
#[cfg(unix)]
impl PidLiveness {
fn is_running(self) -> bool {
!matches!(self, PidLiveness::Dead)
}
}
#[cfg(unix)]
fn classify_kill_result(rc: i32, errno: i32) -> PidLiveness {
if rc == 0 {
return PidLiveness::Alive;
}
match errno {
libc::EPERM => PidLiveness::PermissionDenied,
_ => PidLiveness::Dead,
}
}
#[cfg(unix)]
fn is_process_running(pid: u32) -> bool {
let Ok(pid) = i32::try_from(pid) else {
return false;
};
if pid <= 0 {
return false;
}
let rc = unsafe { libc::kill(pid, 0) };
let errno = std::io::Error::last_os_error().raw_os_error().unwrap_or(0);
classify_kill_result(rc, errno).is_running()
}
#[cfg(unix)]
fn pid_can_name_incumbent(pid: u32, current_pid: u32, allow_same_process_incumbent: bool) -> bool {
allow_same_process_incumbent || pid != current_pid
}
#[cfg(unix)]
const DUPLICATE_PROBE_TIMEOUT: std::time::Duration = std::time::Duration::from_millis(500);
#[cfg(unix)]
async fn socket_speaks_khived_protocol(sock: &std::path::Path, expected_config_id: &str) -> bool {
let probe = DaemonRequestFrame {
probe_only: true,
protocol_version: PROTOCOL_VERSION,
config_id: expected_config_id.to_string(),
..Default::default()
};
let Ok(payload) = serde_json::to_vec(&probe) else {
return false;
};
let response = tokio::time::timeout(DUPLICATE_PROBE_TIMEOUT, async {
let mut stream = UnixStream::connect(sock).await.ok()?;
write_frame(&mut stream, &payload).await.ok()?;
let raw = read_frame(&mut stream).await.ok()?;
serde_json::from_slice::<DaemonResponseFrame>(&raw).ok()
})
.await
.ok()
.flatten();
let Some(resp) = response else {
return false;
};
let is_probe_ack = resp.ok
&& resp.result.is_none()
&& resp.error.is_none()
&& resp.metrics.is_none()
&& resp.request_id.is_none();
is_probe_ack
&& !resp.version_mismatch
&& !resp.namespace_mismatch
&& !resp.config_mismatch
&& resp.daemon_protocol_version == PROTOCOL_VERSION
&& resp
.served_config_id
.as_deref()
.is_some_and(|served| config_ids_compatible(expected_config_id, served))
}
#[cfg(unix)]
async fn socket_is_unreachable(sock: &std::path::Path) -> bool {
match tokio::time::timeout(DUPLICATE_PROBE_TIMEOUT, UnixStream::connect(sock)).await {
Ok(Err(error)) => matches!(
error.kind(),
std::io::ErrorKind::NotFound | std::io::ErrorKind::ConnectionRefused
),
_ => false,
}
}
#[cfg(unix)]
enum Incumbent {
Serving(u32),
Live(u32),
Stale,
}
#[cfg(unix)]
async fn cleanup_stale_daemon(
sock: &std::path::Path,
pid_file: &std::path::Path,
allow_same_process_incumbent: bool,
expected_config_id: &str,
) -> Incumbent {
let mut stale_pid_file_guard = None;
if let Ok(pid_str) = std::fs::read_to_string(pid_file) {
if let Ok(pid) = pid_str.trim().parse::<u32>() {
if pid_can_name_incumbent(pid, std::process::id(), allow_same_process_incumbent)
&& is_process_running(pid)
{
if sock.exists() && socket_speaks_khived_protocol(sock, expected_config_id).await {
return Incumbent::Serving(pid);
}
if sock.exists() && !socket_is_unreachable(sock).await {
return Incumbent::Live(pid);
}
match try_acquire_pid_file_lock(pid_file) {
Ok(Some(guard)) => stale_pid_file_guard = Some(guard),
Ok(None) => return Incumbent::Live(pid),
Err(e) => {
tracing::warn!(
error = %e,
path = ?pid_file,
"cannot check daemon PID-file lock"
);
return Incumbent::Live(pid);
}
}
}
}
}
if sock.exists() {
if let Err(e) = std::fs::remove_file(sock) {
tracing::warn!(error = %e, path = ?sock, "failed to remove stale socket");
}
}
if pid_file.exists() {
if let Err(e) = std::fs::remove_file(pid_file) {
tracing::warn!(error = %e, path = ?pid_file, "failed to remove stale PID file");
}
}
drop(stale_pid_file_guard);
Incumbent::Stale
}
#[cfg(unix)]
fn write_pid_file_exclusive(pid_file: &std::path::Path) -> std::io::Result<std::fs::File> {
use std::os::unix::fs::OpenOptionsExt;
let mut opts = std::fs::OpenOptions::new();
opts.write(true).create_new(true).mode(0o600);
let mut f = opts.open(pid_file)?;
let rc = unsafe { libc::flock(f.as_raw_fd(), libc::LOCK_EX | libc::LOCK_NB) };
if rc != 0 {
return Err(std::io::Error::last_os_error());
}
f.write_all(std::process::id().to_string().as_bytes())?;
Ok(f)
}
#[cfg(unix)]
fn try_acquire_pid_file_lock(pid_file: &std::path::Path) -> std::io::Result<Option<std::fs::File>> {
let file = std::fs::OpenOptions::new()
.read(true)
.write(true)
.open(pid_file)?;
let rc = unsafe { libc::flock(file.as_raw_fd(), libc::LOCK_EX | libc::LOCK_NB) };
if rc == 0 {
return Ok(Some(file));
}
let error = std::io::Error::last_os_error();
if error.kind() == std::io::ErrorKind::WouldBlock
|| error.raw_os_error() == Some(libc::EWOULDBLOCK)
{
Ok(None)
} else {
Err(error)
}
}
#[cfg(unix)]
fn remove_pid_file_if_owned(pid_file: &std::path::Path, guard: &std::fs::File) {
let Ok(owned) = guard.metadata() else {
return;
};
let Ok(current) = std::fs::metadata(pid_file) else {
return;
};
if owned.dev() == current.dev() && owned.ino() == current.ino() {
if let Err(e) = std::fs::remove_file(pid_file) {
tracing::warn!(error = %e, path = ?pid_file, "failed to remove unbound PID file");
}
}
}
#[cfg(unix)]
async fn pid_file_names_a_reachable_daemon(
pid_file: &std::path::Path,
sock: &std::path::Path,
allow_same_process_incumbent: bool,
expected_config_id: &str,
) -> bool {
let Ok(pid_str) = std::fs::read_to_string(pid_file) else {
return false;
};
let Ok(pid) = pid_str.trim().parse::<u32>() else {
return false;
};
pid_can_name_incumbent(pid, std::process::id(), allow_same_process_incumbent)
&& is_process_running(pid)
&& sock.exists()
&& socket_speaks_khived_protocol(sock, expected_config_id).await
}
#[cfg(unix)]
async fn drain(active: &std::sync::atomic::AtomicUsize) -> bool {
drain_with_timeout(active, drain_timeout()).await
}
#[cfg(unix)]
async fn drain_for_idle(active: &std::sync::atomic::AtomicUsize, timeout: std::time::Duration) {
let deadline = tokio::time::Instant::now() + timeout;
let mut warned = false;
while active.load(std::sync::atomic::Ordering::SeqCst) + background_task_count() != 0 {
if !warned && tokio::time::Instant::now() >= deadline {
tracing::warn!(
"idle drain interval elapsed; retaining workers and rendezvous until settled"
);
warned = true;
}
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
}
}
#[cfg(unix)]
async fn drain_with_timeout(
active: &std::sync::atomic::AtomicUsize,
timeout: std::time::Duration,
) -> bool {
use std::sync::atomic::Ordering;
let remaining = || active.load(Ordering::SeqCst) + background_task_count();
if remaining() == 0 {
return true;
}
let deadline = tokio::time::Instant::now() + timeout;
while remaining() > 0 {
if tokio::time::Instant::now() >= deadline {
tracing::warn!(
remaining_connections = active.load(Ordering::SeqCst),
remaining_background_tasks = background_task_count(),
outstanding_background_tasks = %background_task_names().join(", "),
"drain timeout reached; forcing shutdown"
);
return false;
}
tokio::select! {
_ = tokio::time::sleep(std::time::Duration::from_millis(100)) => {}
_ = tokio::time::sleep_until(deadline) => {}
}
}
true
}
#[cfg(unix)]
async fn finish_connection_tasks(tasks: Vec<tokio::task::JoinHandle<()>>, drained: bool) {
if !drained {
for task in &tasks {
if !task.is_finished() {
task.abort();
}
}
}
for task in tasks {
let _ = task.await;
}
}
pub fn drain_timeout() -> std::time::Duration {
let secs = khive_db::env::env_parse_or("KHIVE_DRAIN_TIMEOUT_SECS", DEFAULT_DRAIN_TIMEOUT_SECS);
std::time::Duration::from_secs(secs)
}
#[cfg(unix)]
pub fn env_truthy(key: &str) -> bool {
std::env::var(key)
.map(|v| {
let v = v.trim();
!v.is_empty() && v != "0" && !v.eq_ignore_ascii_case("false")
})
.unwrap_or(false)
}
include!("daemon_khive_root_tests.rs");
#[cfg(all(unix, any(test, feature = "test-internals")))]
#[doc(hidden)]
pub async fn serve_connection_for_test<D: DaemonDispatch>(stream: UnixStream, dispatcher: D) {
handle_conn_with_shutdown(
stream,
dispatcher,
None,
tokio::time::Instant::now() + INITIAL_FRAME_READ_TIMEOUT,
)
.await;
}
#[cfg(all(test, unix))]
#[path = "daemon_tests.rs"]
mod tests;