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 = std::env::var("KHIVE_DRAIN_TIMEOUT_SECS")
.ok()
.and_then(|v| v.parse::<u64>().ok())
.unwrap_or(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))]
mod tests {
include!("daemon/plan_tests.rs");
mod shutdown_signals {
include!("daemon/shutdown_signal_tests.rs");
}
mod connection_limit_tests;
use super::*;
use serial_test::serial;
#[test]
fn lexical_timeout_detail_hides_marker_from_old_clients_without_changing_frame_fit() {
let public = serde_json::json!({
"results": [{"ok": true, "tool": "knowledge.search", "result": "| name |\n|---|\n| first |\n"}],
"summary": {"total": 1, "succeeded": 1, "failed": 0}
});
let public_raw = public.to_string();
let mut marked = public;
marked[DAEMON_LEXICAL_TIMEOUT_MARKER] = serde_json::json!(true);
let marked_raw = marked.to_string();
let (result, detail) = take_daemon_lexical_timeout_marker(marked_raw.clone());
assert_eq!(result, public_raw);
assert_eq!(detail, Some(serde_json::json!({"lexical_timeout": true})));
let frame = |result, error_detail| DaemonResponseFrame {
ok: true,
result: Some(result),
error: None,
error_detail,
namespace_mismatch: false,
config_mismatch: false,
served_config_id: Some("test".to_string()),
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: Some(u64::MAX),
};
let internal_len = serde_json::to_vec(&frame(marked_raw, None)).unwrap().len();
let sent = frame(result, detail);
assert_eq!(sent.result.as_deref(), Some(public_raw.as_str()));
assert!(!sent
.result
.as_deref()
.unwrap()
.contains(DAEMON_LEXICAL_TIMEOUT_MARKER));
assert_eq!(serde_json::to_vec(&sent).unwrap().len(), internal_len);
let untouched = " {\"results\":[],\"summary\":{}} ".to_string();
assert_eq!(
take_daemon_lexical_timeout_marker(untouched.clone()),
(untouched, None)
);
}
#[tokio::test]
async fn incomplete_initial_frames_release_the_connection_deadline() {
for prefix in [&[][..], &[0, 0][..], &[0, 0, 0, 5][..]] {
let (mut peer, mut server) = tokio::io::duplex(64);
peer.write_all(prefix).await.expect("send partial frame");
let deadline = tokio::time::Instant::now() + std::time::Duration::from_millis(10);
let error = read_initial_frame(&mut server, deadline)
.await
.expect_err("an idle peer cannot hold a daemon connection indefinitely");
assert_eq!(error.kind(), std::io::ErrorKind::TimedOut);
}
let (mut peer, mut server) = tokio::io::duplex(64);
write_frame(&mut peer, b"{}")
.await
.expect("send full frame");
assert_eq!(
read_initial_frame(
&mut server,
tokio::time::Instant::now() + std::time::Duration::from_secs(1),
)
.await
.expect("complete frame remains readable"),
b"{}"
);
}
#[test]
fn repeated_accept_failures_back_off_and_cap_at_one_second() {
let mut previous = None;
for expected_ms in [10, 20, 40, 80, 160, 320, 640, 1000, 1000] {
let next = next_accept_error_backoff(previous);
assert_eq!(next.as_millis(), expected_ms);
previous = Some(next);
}
assert_eq!(next_accept_error_backoff(None).as_millis(), 10);
}
#[derive(Debug)]
struct DrainBlockingBlobStore {
started: std::sync::Mutex<Option<tokio::sync::oneshot::Sender<()>>>,
release: Arc<tokio::sync::Semaphore>,
}
#[async_trait]
impl khive_storage::BlobStore for DrainBlockingBlobStore {
async fn put(
&self,
_bytes: Vec<u8>,
) -> khive_storage::StorageResult<khive_storage::ContentRef> {
panic!("put is not used by the hydration drain test")
}
async fn get_bounded_verified(
&self,
_content_ref: &khive_storage::ContentRef,
_max_bytes: u64,
) -> khive_storage::StorageResult<Vec<u8>> {
if let Some(started) = self
.started
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take()
{
let _ = started.send(());
}
self.release
.clone()
.acquire_owned()
.await
.expect("test release semaphore remains open")
.forget();
Ok(b"late result".to_vec())
}
async fn exists(
&self,
_content_ref: &khive_storage::ContentRef,
) -> khive_storage::StorageResult<bool> {
panic!("exists is not used by the hydration drain test")
}
async fn size(
&self,
_content_ref: &khive_storage::ContentRef,
) -> khive_storage::StorageResult<Option<u64>> {
panic!("size is not used by the hydration drain test")
}
async fn delete(
&self,
_content_ref: &khive_storage::ContentRef,
) -> khive_storage::StorageResult<bool> {
panic!("delete is not used by the hydration drain test")
}
}
struct AppendCompletionEventStore {
inner: Arc<dyn khive_storage::EventStore>,
first_append: std::sync::Mutex<Option<tokio::sync::oneshot::Sender<()>>>,
}
#[async_trait]
impl khive_storage::EventStore for AppendCompletionEventStore {
async fn append_event(
&self,
event: khive_storage::Event,
) -> khive_storage::StorageResult<()> {
self.inner.append_event(event).await?;
if let Some(completed) = self
.first_append
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take()
{
let _ = completed.send(());
}
Ok(())
}
async fn append_events(
&self,
events: Vec<khive_storage::Event>,
) -> khive_storage::StorageResult<khive_storage::BatchWriteSummary> {
self.inner.append_events(events).await
}
async fn get_event(
&self,
id: uuid::Uuid,
) -> khive_storage::StorageResult<Option<khive_storage::Event>> {
self.inner.get_event(id).await
}
async fn query_events(
&self,
filter: khive_storage::EventFilter,
page: khive_storage::PageRequest,
) -> khive_storage::StorageResult<khive_storage::Page<khive_storage::Event>> {
self.inner.query_events(filter, page).await
}
async fn count_events(
&self,
filter: khive_storage::EventFilter,
) -> khive_storage::StorageResult<u64> {
self.inner.count_events(filter).await
}
}
#[tokio::test]
#[serial(checkpoint_skip_metrics)]
async fn secondary_only_checkpoint_topology_emits_lifecycle_outcome() {
let main_backend = khive_db::StorageBackend::memory().expect("in-memory main backend");
let inner_event_store = main_backend.events().expect("main event store");
let (first_append_tx, first_append_rx) = tokio::sync::oneshot::channel();
let event_store: Arc<dyn khive_storage::EventStore> =
Arc::new(AppendCompletionEventStore {
inner: inner_event_store,
first_append: std::sync::Mutex::new(Some(first_append_tx)),
});
let secondary_dir = tempfile::tempdir().expect("secondary tempdir");
let secondary_backend =
khive_db::StorageBackend::sqlite_for_test(secondary_dir.path().join("secondary.db"))
.expect("file-backed secondary backend");
let mut tasks = checkpoint_task_specs(
None,
vec![secondary_backend.pool_arc()],
Some(Arc::clone(&event_store)),
"local".to_string(),
);
assert_eq!(tasks.len(), 1);
let task = tasks.pop().expect("one secondary checkpoint task");
assert!(!task.is_main, "the only checkpoint task must be secondary");
assert!(
task.lifecycle_owner.is_some(),
"the secondary task must own lifecycle emission when no main task exists"
);
let config = CheckpointConfig {
interval: std::time::Duration::from_millis(10),
warn_pages: 0,
..CheckpointConfig::default()
};
let (shutdown_tx, shutdown_rx) = tokio::sync::watch::channel(());
let handle = tokio::spawn(run_checkpoint_task(
task.pool,
config,
task.lifecycle_owner,
shutdown_rx,
task.is_main,
));
tokio::time::timeout(std::time::Duration::from_secs(10), first_append_rx)
.await
.expect("secondary checkpoint owner did not complete an append within 10s")
.expect("checkpoint lifecycle append completion sender dropped");
let events = event_store
.query_events(
khive_storage::EventFilter::default(),
khive_storage::PageRequest {
limit: 100,
offset: 0,
},
)
.await
.expect("query lifecycle events");
shutdown_tx.send(()).expect("send checkpoint shutdown");
tokio::time::timeout(std::time::Duration::from_secs(1), handle)
.await
.expect("checkpoint task should exit within 1s")
.expect("checkpoint task panicked");
assert!(
!events.items.is_empty()
&& events
.items
.iter()
.all(|event| event.kind == khive_types::EventKind::CheckpointOutcomeRecorded),
"the designated secondary owner must emit CheckpointOutcomeRecorded"
);
let file_main_dir = tempfile::tempdir().expect("file-backed main tempdir");
let file_main =
khive_db::StorageBackend::sqlite_for_test(file_main_dir.path().join("main.db"))
.expect("file-backed main backend");
let tasks = checkpoint_task_specs(
Some(file_main.pool_arc()),
vec![secondary_backend.pool_arc()],
Some(event_store),
"local".to_string(),
);
assert!(tasks[0].is_main && tasks[0].lifecycle_owner.is_some());
assert!(!tasks[1].is_main && tasks[1].lifecycle_owner.is_none());
}
#[test]
fn current_process_is_running() {
let pid = std::process::id();
assert!(
is_process_running(pid),
"current process {pid} should be detected as running"
);
}
#[test]
fn pid_zero_is_not_running() {
assert!(
!is_process_running(0),
"pid 0 must be rejected by the guard before the unsafe call"
);
}
#[test]
fn very_large_pid_is_not_running() {
assert!(
!is_process_running(u32::MAX),
"u32::MAX should fail i32 conversion and return false"
);
}
#[test]
fn classify_kill_result_zero_is_alive() {
assert_eq!(classify_kill_result(0, 0), PidLiveness::Alive);
assert!(classify_kill_result(0, 0).is_running());
}
#[test]
fn classify_kill_result_esrch_is_dead() {
assert_eq!(classify_kill_result(-1, libc::ESRCH), PidLiveness::Dead);
assert!(!classify_kill_result(-1, libc::ESRCH).is_running());
}
#[test]
fn classify_kill_result_eperm_is_permission_denied_and_counts_as_running() {
assert_eq!(
classify_kill_result(-1, libc::EPERM),
PidLiveness::PermissionDenied
);
assert!(
classify_kill_result(-1, libc::EPERM).is_running(),
"EPERM must be unknown-safe: treated as running, never as a basis \
for stale cleanup to unlink a live daemon's rendezvous files"
);
}
#[test]
fn same_process_pid_requires_explicit_in_process_harness_opt_in() {
let current = std::process::id();
assert!(
!pid_can_name_incumbent(current, current, false),
"production startup must not trust a same-PID stale rendezvous"
);
assert!(
pid_can_name_incumbent(current, current, true),
"the in-process harness must let a live same-PID owner win"
);
let distinct_probe_pid = current.wrapping_add(2);
assert_ne!(
distinct_probe_pid, current,
"probe PID must differ from this process's PID"
);
assert!(
pid_can_name_incumbent(distinct_probe_pid, current, false),
"a distinct PID remains eligible under ordinary production rules"
);
}
#[test]
fn pid_1_probe_is_running_regardless_of_permission_outcome() {
assert!(
is_process_running(1),
"PID 1 always exists; EPERM must not read as dead"
);
}
#[tokio::test]
async fn stale_cleanup_preserves_live_incumbent_without_reachable_socket() {
if crate::test_process::run_in_child() {
return;
}
assert_eq!(
std::env::var("KHIVE_RUNTIME_ISOLATED_TEST").ok().as_deref(),
Some("daemon::tests::stale_cleanup_preserves_live_incumbent_without_reachable_socket"),
"the stale-listener fixture must run alone in its child process"
);
for socket_exists in [false, true] {
let dir = tempfile::tempdir().expect("tempdir");
let sock = dir.path().join("khived.sock");
let pid_file = dir.path().join("khived.pid");
if socket_exists {
let listener = std::os::unix::net::UnixListener::bind(&sock)
.expect("bind socket before closing listener");
drop(listener);
}
let identity = socket_identity(&sock);
assert_eq!(identity.is_some(), socket_exists);
let error = UnixStream::connect(&sock)
.await
.expect_err("incumbent must have no reachable listener");
assert_eq!(
error.kind(),
if socket_exists {
std::io::ErrorKind::ConnectionRefused
} else {
std::io::ErrorKind::NotFound
}
);
let live_pid = std::process::id().to_string();
let _pid_file_guard = write_pid_file_exclusive(&pid_file)
.expect("claim and lock the live incumbent PID file");
assert!(
matches!(
cleanup_stale_daemon(&sock, &pid_file, true, "probe-test").await,
Incumbent::Live(_) | Incumbent::Serving(_)
),
"live incumbent must retain ownership with socket_exists={socket_exists}"
);
assert_eq!(
std::fs::read_to_string(&pid_file).expect("live incumbent PID must survive"),
live_pid
);
assert!(socket_identity(&sock) == identity);
}
}
#[tokio::test]
#[serial]
async fn live_foreign_pid_does_not_block_daemon_startup() {
if crate::test_process::run_in_child() {
return;
}
let dir = tempfile::tempdir().expect("tempdir");
let sock = dir.path().join("khived.sock");
let pid_file = dir.path().join("khived.pid");
std::env::set_var("KHIVE_SOCKET", &sock);
std::env::set_var("KHIVE_PID", &pid_file);
std::env::set_var("KHIVE_LOCK", dir.path().join("khived.recovery.lock"));
let stale_listener =
std::os::unix::net::UnixListener::bind(&sock).expect("create stale socket path");
drop(stale_listener);
let mut foreign = std::process::Command::new("/bin/sleep")
.arg("30")
.spawn()
.expect("spawn live unrelated process");
std::fs::write(&pid_file, foreign.id().to_string()).expect("write unrelated PID");
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "foreign-pid-start-test".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: None,
dispatch_err: None,
};
let daemon = tokio::spawn(run_daemon_in_process_test(dispatcher));
let connected = tokio::time::timeout(std::time::Duration::from_secs(5), async {
loop {
if let Ok(stream) = UnixStream::connect(&sock).await {
break Some(stream);
}
if daemon.is_finished() {
break None;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
})
.await;
let response = if let Ok(Some(mut stream)) = connected {
let mut request = base_request_frame("foreign-pid-start-test");
request.probe_only = true;
let payload = serde_json::to_vec(&request).expect("encode probe request");
tokio::time::timeout(std::time::Duration::from_secs(1), async {
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()
} else {
None
};
let foreign_survived_start = foreign
.try_wait()
.expect("query unrelated process state")
.is_none();
daemon.abort();
let _ = daemon.await;
let _ = foreign.kill();
let _ = foreign.wait();
std::env::remove_var("KHIVE_SOCKET");
std::env::remove_var("KHIVE_PID");
std::env::remove_var("KHIVE_LOCK");
assert!(
response.is_some_and(|response| {
response.ok
&& response.served_config_id.as_deref() == Some("foreign-pid-start-test")
}),
"daemon must start and answer its identity probe"
);
assert!(
foreign_survived_start,
"starting khived must leave the unrelated live process running"
);
}
#[tokio::test]
#[serial]
async fn second_start_refuses_while_pid_file_is_locked_before_bind() {
if crate::test_process::run_in_child() {
return;
}
let dir = tempfile::tempdir().expect("tempdir");
let sock = dir.path().join("khived.sock");
let pid_file = dir.path().join("khived.pid");
std::env::set_var("KHIVE_SOCKET", &sock);
std::env::set_var("KHIVE_PID", &pid_file);
std::env::set_var("KHIVE_LOCK", dir.path().join("khived.recovery.lock"));
let _incumbent_startup_guard = write_pid_file_exclusive(&pid_file)
.expect("incumbent claims and locks its PID file before binding");
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "startup-lock-test".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: None,
dispatch_err: None,
};
let second = tokio::time::timeout(
std::time::Duration::from_secs(1),
run_daemon_in_process_test(dispatcher),
)
.await;
let refused = matches!(second, Ok(Err(_)));
let pid_file_survived = pid_file.exists();
let socket_was_not_bound = !sock.exists();
std::env::remove_var("KHIVE_SOCKET");
std::env::remove_var("KHIVE_PID");
std::env::remove_var("KHIVE_LOCK");
assert!(
refused,
"a second start must refuse while an incumbent holds its pre-bind PID lock"
);
assert!(
pid_file_survived,
"the incumbent PID file must remain in place"
);
assert!(
socket_was_not_bound,
"the second start must not bind the socket"
);
}
#[test]
fn env_truthy_recognises_set_values() {
assert!(!env_truthy("__KHIVE_TEST_ABSENT_VAR_XYZ__"));
let key = "__KHIVE_TEST_TRUTHY_ABC__";
std::env::set_var(key, "1");
assert!(env_truthy(key));
std::env::set_var(key, "false");
assert!(!env_truthy(key));
std::env::set_var(key, "0");
assert!(!env_truthy(key));
std::env::remove_var(key);
}
#[tokio::test(flavor = "current_thread", start_paused = true)]
#[serial(background_tasks)]
async fn accepted_connection_is_counted_before_first_poll_and_drain_waits() {
use std::sync::atomic::Ordering;
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let started = Arc::new(std::sync::atomic::AtomicBool::new(false));
let (release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
let started_in_task = Arc::clone(&started);
let handle = spawn_connection_task(Arc::clone(&active), async move {
started_in_task.store(true, Ordering::Relaxed);
let _ = release_rx.await;
});
assert_eq!(active.load(Ordering::Relaxed), 1);
assert!(
!started.load(Ordering::Relaxed),
"the current-thread runtime must leave the spawned handler unpolled"
);
let drain_fut = drain(active.as_ref());
tokio::pin!(drain_fut);
let too_early =
tokio::time::timeout(std::time::Duration::from_millis(150), &mut drain_fut).await;
assert!(
too_early.is_err(),
"drain must wait for a connection claimed before its task's first poll"
);
assert!(started.load(Ordering::Relaxed));
release_tx.send(()).expect("handler still waiting");
tokio::time::timeout(std::time::Duration::from_secs(1), handle)
.await
.expect("handler should finish promptly")
.expect("handler should not panic");
assert_eq!(active.load(Ordering::Relaxed), 0);
tokio::time::timeout(std::time::Duration::from_secs(1), drain_fut)
.await
.expect("drain should finish once the handler releases its claim");
}
#[tokio::test(flavor = "current_thread")]
async fn cancelled_connection_releases_count_before_first_poll() {
use std::sync::atomic::Ordering;
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let started = Arc::new(std::sync::atomic::AtomicBool::new(false));
let started_in_task = Arc::clone(&started);
let handle = spawn_connection_task(Arc::clone(&active), async move {
started_in_task.store(true, Ordering::Relaxed);
std::future::pending::<()>().await;
});
assert_eq!(active.load(Ordering::Relaxed), 1);
handle.abort();
let error = handle.await.expect_err("aborted handler must be cancelled");
assert!(error.is_cancelled());
assert!(!started.load(Ordering::Relaxed));
assert_eq!(active.load(Ordering::Relaxed), 0);
}
#[tokio::test(flavor = "current_thread")]
async fn panicked_connection_releases_count() {
use std::sync::atomic::Ordering;
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let handle = spawn_connection_task(Arc::clone(&active), async move {
panic!("intentional connection-handler panic");
});
assert_eq!(active.load(Ordering::Relaxed), 1);
let error = handle
.await
.expect_err("panicked handler must fail its join");
assert!(error.is_panic());
assert_eq!(active.load(Ordering::Relaxed), 0);
}
#[test]
fn connection_claim_releases_if_spawn_panics() {
use std::sync::atomic::Ordering;
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
drop(spawn_connection_task(Arc::clone(&active), async {}));
}));
assert!(result.is_err(), "tokio::spawn outside a runtime must panic");
assert_eq!(active.load(Ordering::Relaxed), 0);
}
#[tokio::test]
#[serial(background_tasks)]
async fn drain_returns_promptly_with_no_accepted_connection() {
let active = std::sync::atomic::AtomicUsize::new(0);
tokio::time::timeout(std::time::Duration::from_secs(1), drain(&active))
.await
.expect("empty drain should return immediately");
}
#[test]
fn stopped_listener_is_closed_before_drain() {
use std::process::{Command, Stdio};
let dir = tempfile::Builder::new()
.prefix("kh-drain-")
.tempdir_in("/tmp")
.expect("short isolated socket directory");
let child_home = dir.path().join("home");
std::fs::create_dir(&child_home).expect("empty daemon child HOME");
let mut child = Command::new(std::env::current_exe().expect("test executable"))
.args([
"--exact",
"daemon::tests::stopped_listener_is_closed_before_drain_child",
"--ignored",
"--nocapture",
"--test-threads=1",
])
.env_clear()
.envs(
std::env::vars_os().filter(|(key, _)| !key.to_string_lossy().starts_with("KHIVE_")),
)
.env("HOME", &child_home)
.env("KHIVE_VOLUME_LOCK_DIR", dir.path().join("volume-locks"))
.env_remove("LATTICE_MODEL_CACHE")
.env("KHIVE_TEST_HARNESS", "1")
.env("KHIVE_DRAIN_TEST_CHILD", "1")
.env("KHIVE_SOCKET", dir.path().join("s"))
.env("KHIVE_PID", dir.path().join("p"))
.env("KHIVE_LOCK", dir.path().join("l"))
.env("KHIVE_RECOVERER_LOCK", dir.path().join("r"))
.env("KHIVE_DRAIN_TIMEOUT_SECS", "10")
.current_dir(dir.path())
.stdin(Stdio::null())
.stdout(Stdio::piped())
.stderr(Stdio::piped())
.spawn()
.expect("spawn isolated daemon test");
let deadline = std::time::Instant::now() + std::time::Duration::from_secs(15);
let completed = loop {
match child.try_wait() {
Ok(Some(_)) => break true,
Ok(None) if std::time::Instant::now() < deadline => {
std::thread::sleep(std::time::Duration::from_millis(10));
}
_ => {
let _ = child.kill();
break false;
}
}
};
let output = child.wait_with_output().expect("reap daemon test child");
assert!(completed, "daemon test child did not finish: {output:?}");
assert!(output.status.success(), "daemon test failed: {output:?}");
assert!(
String::from_utf8_lossy(&output.stdout).contains("STOPPED_LISTENER_DRAIN_VERIFIED"),
"child must run the listener witness: {output:?}"
);
assert!(
std::fs::read_dir(child_home).unwrap().next().is_none(),
"daemon drain child must leave its private HOME empty"
);
}
#[tokio::test]
#[ignore = "subprocess helper, invoked by stopped_listener_is_closed_before_drain"]
async fn stopped_listener_is_closed_before_drain_child() {
assert_eq!(
std::env::var("KHIVE_DRAIN_TEST_CHILD").expect("isolated child environment"),
"1"
);
let _sigterm = tokio::signal::unix::signal(tokio::signal::unix::SignalKind::terminate())
.expect("install child SIGTERM handler");
let (release_tx, release_rx) = tokio::sync::oneshot::channel::<()>();
let background = spawn_tracked_task(async move {
release_rx.await.expect("release held drain task");
});
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "drain-test".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: None,
dispatch_err: None,
};
let starts = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let stopped = Arc::new(std::sync::atomic::AtomicBool::new(false));
let callback_starts = Arc::clone(&starts);
let callback_stopped = Arc::clone(&stopped);
let boot_guard = Some(acquire_daemon_boot_guard().expect("boot guard"));
let daemon = tokio::spawn(run_daemon_with_boot_guard_and_start(
dispatcher,
boot_guard,
move |_| {
use std::os::unix::fs::FileTypeExt;
assert!(std::fs::metadata(socket_path())
.unwrap()
.file_type()
.is_socket());
assert_eq!(
std::fs::read_to_string(pid_path()).unwrap(),
std::process::id().to_string()
);
assert_eq!(
callback_starts.fetch_add(1, std::sync::atomic::Ordering::SeqCst),
0
);
track_named_background_task("startup-lifecycle-test", async move {
daemon_shutdown_token().cancelled().await;
callback_stopped.store(true, std::sync::atomic::Ordering::SeqCst);
});
},
));
let sock = socket_path();
let mut stream = tokio::time::timeout(std::time::Duration::from_secs(5), async {
loop {
if let Ok(stream) = UnixStream::connect(&sock).await {
break stream;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
})
.await
.expect("daemon must bind");
let payload = serde_json::to_vec(&base_request_frame("drain-test"))
.expect("encode readiness request");
write_frame(&mut stream, &payload)
.await
.expect("write readiness request");
let response =
tokio::time::timeout(std::time::Duration::from_secs(2), read_frame(&mut stream))
.await
.expect("daemon must serve readiness request")
.expect("read readiness response");
let response: DaemonResponseFrame =
serde_json::from_slice(&response).expect("decode readiness response");
assert!(response.ok, "daemon readiness failed: {response:?}");
assert_eq!(starts.load(std::sync::atomic::Ordering::SeqCst), 1);
drop(stream);
let rc = unsafe { libc::kill(std::process::id() as i32, libc::SIGTERM) };
assert_eq!(rc, 0, "signal isolated daemon child");
tokio::time::timeout(
std::time::Duration::from_secs(2),
daemon_shutdown_token().cancelled(),
)
.await
.expect("daemon must begin shutdown");
assert!(
!daemon.is_finished(),
"held background task must retain drain"
);
assert!(
sock.exists(),
"cleanup must not have removed the socket yet"
);
assert_eq!(
std::fs::read_to_string(pid_path()).expect("draining daemon PID"),
std::process::id().to_string()
);
let late_connect = tokio::time::timeout(
std::time::Duration::from_secs(1),
UnixStream::connect(&sock),
)
.await
.expect("late connect must finish promptly");
release_tx.send(()).expect("release daemon drain");
background.await.expect("held background task must finish");
tokio::time::timeout(std::time::Duration::from_secs(2), daemon)
.await
.expect("released daemon must finish shutdown")
.expect("daemon task must not panic")
.expect("daemon shutdown must succeed");
let error = late_connect.expect_err("stopped listener must not queue new connections");
assert_eq!(error.kind(), std::io::ErrorKind::ConnectionRefused);
assert!(!sock.exists(), "owned socket must be removed after drain");
assert!(
!pid_path().exists(),
"owned PID must be removed after drain"
);
assert!(
stopped.load(std::sync::atomic::Ordering::SeqCst),
"work started after ownership must finish inside daemon drain"
);
println!("STOPPED_LISTENER_DRAIN_VERIFIED");
}
#[tokio::test(start_paused = true)]
#[serial(background_tasks)]
async fn graceful_drain_has_a_hard_upper_bound_with_stuck_work() {
use std::sync::atomic::Ordering;
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let task = spawn_connection_task(Arc::clone(&active), async {
std::future::pending::<()>().await;
});
assert_eq!(active.load(Ordering::Relaxed), 1);
let started = tokio::time::Instant::now();
let drained = drain_with_timeout(&active, std::time::Duration::from_millis(250)).await;
assert!(!drained, "stuck work must exhaust the drain bound");
assert!(
started.elapsed() >= std::time::Duration::from_millis(250)
&& started.elapsed() < std::time::Duration::from_millis(350),
"graceful shutdown exceeded its configured bound: {:?}",
started.elapsed()
);
finish_connection_tasks(vec![task], drained).await;
assert_eq!(
active.load(Ordering::Relaxed),
0,
"hard-bound escalation must abort, await, and release the handler"
);
}
#[tokio::test(start_paused = true)]
#[serial(background_tasks)]
async fn admitted_work_finishes_inside_drain_window_without_abort() {
use std::sync::atomic::Ordering;
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let committed = Arc::new(std::sync::atomic::AtomicBool::new(false));
let committed_in_task = Arc::clone(&committed);
let task = spawn_connection_task(Arc::clone(&active), async move {
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
committed_in_task.store(true, Ordering::SeqCst);
});
let drained = drain_with_timeout(&active, std::time::Duration::from_millis(250)).await;
assert!(
drained,
"admitted work should finish inside the drain window"
);
finish_connection_tasks(vec![task], drained).await;
assert!(committed.load(Ordering::SeqCst));
assert_eq!(active.load(Ordering::Relaxed), 0);
}
#[tokio::test]
#[serial(background_tasks)]
async fn drain_waits_for_tracked_background_tasks_before_returning() {
let active = std::sync::atomic::AtomicUsize::new(0);
let (tx, rx) = tokio::sync::oneshot::channel::<()>();
track_background_task(async move {
let _ = rx.await;
});
assert!(
background_task_count() >= 1,
"track_background_task must make the in-flight task visible immediately"
);
let drain_fut = drain(&active);
tokio::pin!(drain_fut);
let too_early =
tokio::time::timeout(std::time::Duration::from_millis(150), &mut drain_fut).await;
assert!(
too_early.is_err(),
"drain() must not return while a tracked background task is still running"
);
tx.send(())
.expect("tracked task still awaiting the oneshot");
let done = tokio::time::timeout(std::time::Duration::from_secs(5), drain_fut).await;
assert!(
done.is_ok(),
"drain() must return once the tracked background task finishes"
);
}
#[tokio::test]
#[serial(background_tasks)]
async fn drain_waits_for_hydration_after_its_last_request_waiter_is_cancelled() {
let before = background_task_count();
let active = std::sync::atomic::AtomicUsize::new(0);
let (started_tx, started_rx) = tokio::sync::oneshot::channel();
let release = Arc::new(tokio::sync::Semaphore::new(0));
let store = Arc::new(DrainBlockingBlobStore {
started: std::sync::Mutex::new(Some(started_tx)),
release: Arc::clone(&release),
});
let hydrator = Arc::new(
crate::BlobHydrator::new(
store as Arc<dyn khive_storage::BlobStore>,
khive_storage::MAX_BLOB_WHOLE_BYTES,
)
.expect("minimum hydration budget is valid"),
);
let content_ref =
khive_storage::ContentRef::from_hex("a".repeat(64)).expect("fixture content ref");
let request_hydrator = Arc::clone(&hydrator);
let request = tokio::spawn(async move {
request_hydrator
.hydrate_verified(&content_ref, khive_storage::MAX_BLOB_WHOLE_BYTES)
.await
});
started_rx.await.expect("backend work must begin");
request.abort();
assert!(request.await.unwrap_err().is_cancelled());
assert_eq!(background_task_count(), before + 1);
let draining = drain_with_timeout(&active, std::time::Duration::from_secs(5));
tokio::pin!(draining);
assert!(
tokio::time::timeout(std::time::Duration::from_millis(150), &mut draining)
.await
.is_err(),
"drain must remain pending while cancelled-request hydration still runs"
);
release.add_permits(1);
assert!(
tokio::time::timeout(std::time::Duration::from_secs(5), draining)
.await
.expect("drain should finish after native hydration ends"),
"hydration should finish inside the drain window"
);
assert_eq!(background_task_count(), before);
}
#[tokio::test]
#[serial(background_tasks)]
async fn track_background_task_count_returns_to_zero_after_completion() {
let before = background_task_count();
let (tx, rx) = tokio::sync::oneshot::channel::<()>();
track_background_task(async move {
let _ = rx.await;
});
assert_eq!(background_task_count(), before + 1);
tx.send(()).expect("still awaiting");
for _ in 0..100 {
if background_task_count() == before {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
assert_eq!(background_task_count(), before);
}
#[tokio::test]
#[serial(background_tasks)]
async fn track_background_task_count_returns_to_baseline_after_panic() {
let before = background_task_count();
let (tx, rx) = tokio::sync::oneshot::channel::<()>();
track_background_task(async move {
let _ = rx.await;
panic!("intentional panic to exercise the Drop-guard decrement path");
});
assert_eq!(background_task_count(), before + 1);
tx.send(()).expect("still awaiting");
for _ in 0..100 {
if background_task_count() == before {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
assert_eq!(
background_task_count(),
before,
"background task counter must return to baseline after the tracked future panics"
);
}
#[test]
#[serial(active_phases)]
fn register_active_phase_appears_and_disappears_with_the_guard() {
assert!(
!active_phase_names().contains(&"adr103_test_phase".to_string()),
"must start absent (leaked from a prior failed run would poison this test)"
);
let guard = register_active_phase("adr103_test_phase");
assert!(active_phase_names().contains(&"adr103_test_phase".to_string()));
drop(guard);
assert!(
!active_phase_names().contains(&"adr103_test_phase".to_string()),
"the phase name must drop out of the gauge once its guard is dropped"
);
}
#[test]
#[serial(active_phases)]
fn register_active_phase_counts_concurrent_occurrences_of_the_same_name() {
let first = register_active_phase("adr103_concurrent_phase");
let second = register_active_phase("adr103_concurrent_phase");
assert!(active_phase_names().contains(&"adr103_concurrent_phase".to_string()));
drop(first);
assert!(
active_phase_names().contains(&"adr103_concurrent_phase".to_string()),
"one of two concurrent occurrences ending must not remove the name early"
);
drop(second);
assert!(
!active_phase_names().contains(&"adr103_concurrent_phase".to_string()),
"the name must be removed only once every concurrent occurrence has ended"
);
}
#[derive(Clone)]
struct MockDispatch {
namespace: String,
config_id: String,
dispatch_calls: Arc<std::sync::atomic::AtomicUsize>,
pool: Option<Arc<ConnectionPool>>,
dispatch_err: Option<String>,
}
#[derive(Clone)]
struct CancellationAwareDispatch {
started: Arc<tokio::sync::Notify>,
cancellation_observed: Arc<std::sync::atomic::AtomicBool>,
count_sql: Option<Arc<dyn khive_storage::SqlAccess>>,
}
#[async_trait]
impl DaemonDispatch for CancellationAwareDispatch {
fn plan(&self, ops: &str) -> String {
khive_request::plan_request(ops, &Default::default()).to_string()
}
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> {
self.started.notify_one();
if let Some(sql) = &self.count_sql {
let mut reader = sql.reader().await.map_err(|error| error.to_string())?;
let result = reader.query_scalar(khive_storage::SqlStatement {
sql: "SELECT COUNT(*) FROM events WHERE namespace = ?1 AND verb LIKE 'knowledge.%'".into(),
params: vec![khive_storage::SqlValue::Text("local".into())],
label: Some("knowledge.stats.event_count".into()),
}).await;
self.cancellation_observed.store(
matches!(
result,
Err(khive_storage::error::StorageError::Timeout { .. })
),
std::sync::atomic::Ordering::SeqCst,
);
return result
.map(|value| format!("{value:?}"))
.map_err(|error| error.to_string());
}
khive_storage::wait_for_request_read_cancellation().await;
self.cancellation_observed
.store(true, std::sync::atomic::Ordering::SeqCst);
Ok("{}".to_string())
}
async fn warm_all(&self) {}
fn namespace(&self) -> &str {
"local"
}
fn config_id(&self) -> &str {
"disconnect-test"
}
}
mod demand_retirement_tests {
use super::*;
use khive_storage::SqlAccess;
pub(super) fn dispatcher(pool: Option<Arc<ConnectionPool>>) -> MockDispatch {
MockDispatch {
namespace: "local".to_owned(),
config_id: "idle-test".to_owned(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool,
dispatch_err: None,
}
}
pub(super) fn lifecycle(mode: DaemonLifetime) -> Arc<DaemonLifecycle> {
Arc::new(DaemonLifecycle::new(
DaemonOptions {
lifetime: mode,
idle_interval: std::time::Duration::from_secs(1),
},
DaemonStartupReport::default(),
))
}
#[tokio::test(start_paused = true)]
async fn ordinary_cleanup_resets_idle_and_admission_is_one_way() {
let state = lifecycle(DaemonLifetime::Demand);
tokio::time::advance(std::time::Duration::from_secs(3)).await;
assert!(
!state.try_idle(Vec::new),
"readiness must precede the idle clock"
);
state.ready();
let request = state.admit().unwrap();
tokio::time::advance(std::time::Duration::from_secs(3)).await;
assert!(
!state.try_idle(Vec::new),
"admitted work must prevent retirement"
);
drop(request);
assert!(
!state.try_idle(Vec::new),
"cleanup starts a fresh idle interval"
);
tokio::time::advance(std::time::Duration::from_secs(1)).await;
assert!(state.try_idle(Vec::new));
assert!(
state.admit().is_none(),
"draining must refuse before dispatch"
);
assert!(!state.try_idle(Vec::new), "retirement cannot be repeated");
assert_eq!(
state.snapshot().shutdown_reason,
Some(DaemonShutdownReason::Idle)
);
state.stopped();
assert!(state.admit().is_none());
}
#[test]
fn concurrent_admission_and_idle_decision_choose_one_winner() {
for _ in 0..16 {
let state = Arc::new(DaemonLifecycle::new(
DaemonOptions {
lifetime: DaemonLifetime::Demand,
idle_interval: std::time::Duration::from_nanos(1),
},
DaemonStartupReport::default(),
));
state.ready();
let barrier = Arc::new(std::sync::Barrier::new(2));
let admitting_state = Arc::clone(&state);
let admitting_barrier = Arc::clone(&barrier);
let admission = std::thread::spawn(move || {
admitting_barrier.wait();
admitting_state.admit()
});
let retiring_state = Arc::clone(&state);
let retirement = std::thread::spawn(move || {
barrier.wait();
retiring_state.try_idle(Vec::new)
});
let admitted = admission.join().unwrap();
let retired = retirement.join().unwrap();
assert_ne!(
admitted.is_some(),
retired,
"request admission and voluntary retirement cannot both win"
);
if retired {
assert!(state.admit().is_none());
}
drop(admitted);
}
}
#[tokio::test(start_paused = true)]
async fn named_service_obligations_and_unknown_resources_are_ineligible() {
let state = Arc::new(DaemonLifecycle::new(
DaemonOptions {
lifetime: DaemonLifetime::Demand,
idle_interval: std::time::Duration::from_secs(1),
},
DaemonStartupReport {
skipped_components: vec!["schedule-tick".to_owned()],
idle_ineligible_reasons: vec![
"unclassified_component:external-service".to_owned()
],
},
));
state.ready();
tokio::time::advance(std::time::Duration::from_secs(3)).await;
assert!(!state.try_idle(Vec::new));
assert_eq!(
state.snapshot().idle_ineligible_reasons,
vec!["unclassified_component:external-service"]
);
let unknown = CancellationAwareDispatch {
started: Arc::new(tokio::sync::Notify::new()),
cancellation_observed: Arc::new(std::sync::atomic::AtomicBool::new(false)),
count_sql: None,
};
let clean = lifecycle(DaemonLifetime::Demand);
clean.ready();
tokio::time::advance(std::time::Duration::from_secs(3)).await;
assert!(!clean.try_idle(|| unknown.idle_retirement_blockers()));
assert_eq!(
clean.snapshot().idle_blockers,
vec!["dispatcher_resource_inventory_unknown"]
);
}
#[tokio::test(start_paused = true)]
#[serial(background_tasks, tx_registry)]
async fn retained_raw_sql_writer_blocks_actual_idle_wait_and_persistent_stays() {
let dir = tempfile::tempdir().unwrap();
let pool = Arc::new(
ConnectionPool::new(khive_db::PoolConfig {
path: Some(dir.path().join("retained.db")),
write_queue_enabled: Some(false),
write_routing_strict: false,
..Default::default()
})
.unwrap(),
);
let bridge = khive_db::SqlBridge::new(Arc::clone(&pool), true);
let writer = bridge.writer().await.unwrap();
assert!(
khive_storage::tx_registry::snapshot().is_empty(),
"this hold must be autocommit, not an open transaction"
);
let d = dispatcher(Some(Arc::clone(&pool)));
let demand = lifecycle(DaemonLifetime::Demand);
let persistent = lifecycle(DaemonLifetime::Persistent);
demand.ready();
persistent.ready();
tokio::time::advance(std::time::Duration::from_secs(3)).await;
let idle = wait_for_idle(&d, &demand);
tokio::pin!(idle);
assert!(
tokio::time::timeout(std::time::Duration::from_millis(150), &mut idle)
.await
.is_err(),
"a genuine retained writer handle must prevent the actual idle arm"
);
assert!(tokio::time::timeout(
std::time::Duration::from_millis(150),
wait_for_idle(&d, &persistent)
)
.await
.is_err());
drop(writer);
assert_eq!(pool.retirement_writer_holds(), 0);
tokio::time::timeout(std::time::Duration::from_secs(2), idle)
.await
.unwrap();
assert_eq!(demand.snapshot().phase, DaemonLifecyclePhase::Draining);
assert_eq!(persistent.snapshot().phase, DaemonLifecyclePhase::Serving);
let pooled = lifecycle(DaemonLifetime::Demand);
pooled.ready();
tokio::time::advance(std::time::Duration::from_secs(2)).await;
let pooled_guard = pool.writer().unwrap();
assert!(
!pooled.try_idle(|| idle_retirement_blockers(&d)),
"pooled writer hold must block retirement"
);
drop(pooled_guard);
assert!(pooled.try_idle(|| idle_retirement_blockers(&d)));
}
#[tokio::test(start_paused = true)]
#[serial(background_tasks, tx_registry)]
async fn persistent_idle_wait_never_retires_after_writer_release() {
let dir = tempfile::tempdir().unwrap();
let pool = Arc::new(
ConnectionPool::new(khive_db::PoolConfig {
path: Some(dir.path().join("persistent.db")),
write_queue_enabled: Some(false),
write_routing_strict: false,
..Default::default()
})
.unwrap(),
);
let bridge = khive_db::SqlBridge::new(Arc::clone(&pool), true);
let held = bridge.writer().await.unwrap();
let d = dispatcher(Some(pool));
let state = lifecycle(DaemonLifetime::Persistent);
state.ready();
tokio::time::advance(std::time::Duration::from_secs(3)).await;
assert!(tokio::time::timeout(
std::time::Duration::from_millis(100),
wait_for_idle(&d, &state)
)
.await
.is_err());
drop(held);
assert!(tokio::time::timeout(
std::time::Duration::from_secs(2),
wait_for_idle(&d, &state)
)
.await
.is_err());
assert_eq!(state.snapshot().phase, DaemonLifecyclePhase::Serving);
}
#[tokio::test(start_paused = true)]
#[serial(background_tasks, tx_registry)]
async fn explicit_sql_reader_transaction_blocks_retirement_without_writer_hold() {
let dir = tempfile::tempdir().unwrap();
let pool = Arc::new(
ConnectionPool::new(khive_db::PoolConfig {
path: Some(dir.path().join("reader.db")),
write_queue_enabled: Some(false),
write_routing_strict: false,
..Default::default()
})
.unwrap(),
);
let bridge = khive_db::SqlBridge::new(Arc::clone(&pool), true);
let mut reader = bridge.reader().await.unwrap();
reader
.query_all(khive_storage::SqlStatement {
sql: "BEGIN DEFERRED".to_owned(),
params: vec![],
label: Some("idle-reader".to_owned()),
})
.await
.unwrap();
assert_eq!(pool.retirement_writer_holds(), 0);
let d = dispatcher(Some(pool));
let state = lifecycle(DaemonLifetime::Demand);
state.ready();
tokio::time::advance(std::time::Duration::from_secs(3)).await;
assert!(!state.try_idle(|| idle_retirement_blockers(&d)));
assert!(state
.snapshot()
.idle_blockers
.contains(&"open_sql_transaction".to_owned()));
drop(reader);
assert!(state.try_idle(|| idle_retirement_blockers(&d)));
}
#[tokio::test(start_paused = true)]
#[serial(background_tasks, tx_registry)]
async fn unsettled_named_worker_blocks_idle_without_resetting_clock() {
let (release, pending) = tokio::sync::oneshot::channel::<()>();
let task = spawn_named_tracked_task("idle-test-worker", async move {
pending.await.unwrap();
});
let state = lifecycle(DaemonLifetime::Demand);
state.ready();
let d = dispatcher(None);
tokio::time::advance(std::time::Duration::from_secs(2)).await;
assert!(!state.try_idle(|| idle_retirement_blockers(&d)));
assert!(state
.snapshot()
.idle_blockers
.contains(&"unsettled_worker:idle-test-worker".to_owned()));
release.send(()).unwrap();
task.await.unwrap();
assert!(
state.try_idle(|| idle_retirement_blockers(&d)),
"maintenance completion must not reset ordinary activity"
);
}
#[tokio::test(start_paused = true)]
#[serial(background_tasks)]
async fn voluntary_drain_retains_pending_work_past_deadline() {
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let (release, pending) = tokio::sync::oneshot::channel::<()>();
let task = spawn_connection_task(Arc::clone(&active), async move {
pending.await.unwrap();
});
let drain = drain_for_idle(&active, std::time::Duration::from_millis(10));
tokio::pin!(drain);
assert!(
tokio::time::timeout(std::time::Duration::from_secs(1), &mut drain)
.await
.is_err(),
"voluntary timeout must retain admitted work"
);
assert!(!task.is_finished());
release.send(()).unwrap();
task.await.unwrap();
tokio::time::timeout(std::time::Duration::from_secs(1), drain)
.await
.unwrap();
}
#[tokio::test(start_paused = true)]
async fn stalled_response_transport_is_bounded() {
let (mut writer, _held_reader) = tokio::io::duplex(1);
let error = tokio::time::timeout(
std::time::Duration::from_secs(35),
write_response_frame(&mut writer, b"bounded response"),
)
.await
.expect("the production response bound must fire before the fixture ceiling")
.unwrap_err();
assert_eq!(error.kind(), std::io::ErrorKind::TimedOut);
}
#[tokio::test(start_paused = true)]
async fn draining_handler_refuses_before_dispatch() {
let d = dispatcher(None);
let calls = Arc::clone(&d.dispatch_calls);
let state = lifecycle(DaemonLifetime::Demand);
state.ready();
tokio::time::advance(std::time::Duration::from_secs(2)).await;
assert!(state.try_idle(Vec::new));
tokio::time::resume();
let (mut client, server) = UnixStream::pair().unwrap();
let handle = tokio::spawn(handle_conn_with_lifecycle(
server,
d,
None,
tokio::time::Instant::now() + INITIAL_FRAME_READ_TIMEOUT,
Some(state),
));
let mut frame = base_request_frame("idle-test");
frame.ops = "stats()".to_owned();
write_frame(&mut client, &serde_json::to_vec(&frame).unwrap())
.await
.unwrap();
let refusal: DaemonResponseFrame =
serde_json::from_slice(&read_frame(&mut client).await.unwrap()).unwrap();
assert!(!refusal.ok);
assert_eq!(refusal.error_detail.unwrap()["code"], "daemon_draining");
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 0);
handle.await.unwrap();
}
#[test]
fn lifecycle_metrics_are_additive_and_generation_is_stable() {
let state = lifecycle(DaemonLifetime::Demand);
let generation = state.snapshot().instance_generation;
let metrics = MetricsSnapshot {
lifecycle: Some(state.snapshot()),
..Default::default()
};
let decoded: MetricsSnapshot =
serde_json::from_value(serde_json::to_value(metrics).unwrap()).unwrap();
assert_eq!(decoded.lifecycle.unwrap().instance_generation, generation);
let old = serde_json::to_value(MetricsSnapshot::default()).unwrap();
assert!(old.get("lifecycle").is_none());
assert!(serde_json::from_value::<MetricsSnapshot>(old)
.unwrap()
.lifecycle
.is_none());
}
}
#[async_trait]
impl DaemonDispatch for MockDispatch {
fn idle_retirement_blockers(&self) -> Vec<String> {
self.pool
.as_ref()
.filter(|pool| pool.retirement_writer_holds() != 0)
.map(|_| vec!["test_backend:held_writer".to_owned()])
.unwrap_or_default()
}
fn plan(&self, ops: &str) -> String {
khive_request::plan_request(ops, &Default::default()).to_string()
}
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> {
self.dispatch_calls
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
match &self.dispatch_err {
Some(msg) => Err(msg.clone()),
None => Ok("{}".to_string()),
}
}
async fn warm_all(&self) {}
fn namespace(&self) -> &str {
&self.namespace
}
fn config_id(&self) -> &str {
&self.config_id
}
fn pool_for_checkpoint(&self) -> Option<Arc<ConnectionPool>> {
self.pool.clone()
}
}
fn base_request_frame(config_id: &str) -> DaemonRequestFrame {
DaemonRequestFrame {
plan: false,
ops: String::new(),
presentation: None,
presentation_per_op: None,
namespace: "local".to_string(),
actor_id: None,
process_ref: None,
visible_namespaces: Vec::new(),
config_id: config_id.to_string(),
protocol_version: PROTOCOL_VERSION,
probe_only: false,
metrics_only: false,
format: None,
format_per_op: None,
from_wire: false,
request_id: None,
}
}
async fn round_trip<D: DaemonDispatch>(
dispatcher: D,
req: &DaemonRequestFrame,
) -> DaemonResponseFrame {
let (mut client, server) = UnixStream::pair().expect("unix stream pair");
let payload = serde_json::to_vec(req).expect("encode request frame");
let handle = tokio::spawn(async move {
handle_conn(server, dispatcher).await;
});
write_frame(&mut client, &payload)
.await
.expect("write request frame");
let raw = read_frame(&mut client).await.expect("read response frame");
handle.await.expect("handle_conn task panicked");
serde_json::from_slice(&raw).expect("decode response frame")
}
#[tokio::test]
async fn expired_accepted_deadline_refuses_even_buffered_complete_frame() {
let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let dispatcher = MockDispatch {
namespace: "local".into(),
config_id: "expired-accept-test".into(),
dispatch_calls: Arc::clone(&calls),
pool: None,
dispatch_err: None,
};
let (mut client, server) = UnixStream::pair().expect("unix stream pair");
let frame = base_request_frame("expired-accept-test");
write_frame(&mut client, &serde_json::to_vec(&frame).unwrap())
.await
.expect("buffer complete frame before handler starts");
let accepted_deadline = tokio::time::Instant::now() - std::time::Duration::from_secs(1);
tokio::time::timeout(
std::time::Duration::from_secs(1),
handle_conn_with_shutdown(server, dispatcher, None, accepted_deadline),
)
.await
.expect("expired accepted deadline must not start a fresh read window");
let mut byte = [0u8; 1];
match client.read(&mut byte).await {
Ok(0) => {}
Err(error) if error.kind() == std::io::ErrorKind::ConnectionReset => {}
other => panic!("expected closed socket, got {other:?}"),
}
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 0);
}
struct ReadyFrameReader {
frame: Vec<u8>,
offset: usize,
polls: usize,
}
impl tokio::io::AsyncRead for ReadyFrameReader {
fn poll_read(
self: std::pin::Pin<&mut Self>,
_cx: &mut std::task::Context<'_>,
buf: &mut tokio::io::ReadBuf<'_>,
) -> std::task::Poll<std::io::Result<()>> {
let reader = self.get_mut();
reader.polls += 1;
let remaining = &reader.frame[reader.offset..];
let count = remaining.len().min(buf.remaining());
buf.put_slice(&remaining[..count]);
reader.offset += count;
std::task::Poll::Ready(Ok(()))
}
}
#[tokio::test]
async fn expired_accepted_deadline_refuses_a_frame_ready_on_first_poll() {
let mut reader = ReadyFrameReader {
frame: [2_u32.to_be_bytes().as_slice(), b"{}"].concat(),
offset: 0,
polls: 0,
};
let accepted_deadline = tokio::time::Instant::now() - std::time::Duration::from_secs(1);
let error = read_initial_frame(&mut reader, accepted_deadline)
.await
.expect_err("a fully ready frame must not outlive its acceptance deadline");
assert_eq!(error.kind(), std::io::ErrorKind::TimedOut);
assert_eq!(reader.polls, 0, "an expired frame must not be polled");
}
#[tokio::test]
async fn socket_speaks_khived_protocol_accepts_a_real_khived() {
let dir = tempfile::tempdir().expect("tempdir");
let sock_path = dir.path().join("real.sock");
let listener = UnixListener::bind(&sock_path).expect("bind real listener");
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "probe-test".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: None,
dispatch_err: None,
};
let accept_task = tokio::spawn(async move {
if let Ok((stream, _)) = listener.accept().await {
handle_conn(stream, dispatcher).await;
}
});
assert!(
socket_speaks_khived_protocol(&sock_path, "probe-test").await,
"a real khived answering the probe_only frame with a matching config_id must be recognized"
);
let _ = tokio::time::timeout(std::time::Duration::from_secs(2), accept_task).await;
}
include!("daemon/probe_listener_tests.rs");
#[tokio::test]
async fn socket_speaks_khived_protocol_rejects_a_non_ack_or_mismatched_response() {
let dir = tempfile::tempdir().expect("tempdir");
let sock_path = dir.path().join("mismatched.sock");
let listener = UnixListener::bind(&sock_path).expect("bind fake listener");
let accept_task = tokio::spawn(async move {
if let Ok((mut stream, _)) = listener.accept().await {
let _raw = read_frame(&mut stream).await.expect("read probe frame");
let resp = DaemonResponseFrame {
ok: false,
result: None,
error: None,
namespace_mismatch: false,
config_mismatch: true,
served_config_id: Some("someone-elses-config".to_string()),
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
error_detail: None,
metrics: None,
request_id: None,
};
let payload = serde_json::to_vec(&resp).expect("encode response");
write_frame(&mut stream, &payload)
.await
.expect("write response");
}
});
let speaks = socket_speaks_khived_protocol(&sock_path, "expected-config").await;
assert!(
!speaks,
"a well-formed but non-ack / identity-mismatched response must not be treated as \
the same live khived"
);
let _ = tokio::time::timeout(std::time::Duration::from_secs(2), accept_task).await;
}
#[tokio::test]
async fn socket_speaks_khived_protocol_rejects_a_metrics_only_response() {
let dir = tempfile::tempdir().expect("tempdir");
let sock_path = dir.path().join("metrics-only.sock");
let listener = UnixListener::bind(&sock_path).expect("bind fake listener");
let accept_task = tokio::spawn(async move {
if let Ok((mut stream, _)) = listener.accept().await {
let _raw = read_frame(&mut stream).await.expect("read probe frame");
let resp = DaemonResponseFrame {
ok: true,
result: None,
error: None,
namespace_mismatch: false,
config_mismatch: false,
served_config_id: Some("expected-config".to_string()),
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
error_detail: None,
metrics: Some(MetricsSnapshot::default()),
request_id: None,
};
let payload = serde_json::to_vec(&resp).expect("encode response");
write_frame(&mut stream, &payload)
.await
.expect("write response");
}
});
let speaks = socket_speaks_khived_protocol(&sock_path, "expected-config").await;
assert!(
!speaks,
"an otherwise-matching response carrying a metrics snapshot must not be treated as \
a probe acknowledgement"
);
tokio::time::timeout(std::time::Duration::from_secs(2), accept_task)
.await
.expect("fake listener accept task timed out")
.expect("fake listener accept task panicked");
}
#[tokio::test]
async fn socket_speaks_khived_protocol_rejects_a_response_with_request_id() {
let dir = tempfile::tempdir().expect("tempdir");
let sock_path = dir.path().join("request-id.sock");
let listener = UnixListener::bind(&sock_path).expect("bind fake listener");
let accept_task = tokio::spawn(async move {
if let Ok((mut stream, _)) = listener.accept().await {
let _raw = read_frame(&mut stream).await.expect("read probe frame");
let resp = DaemonResponseFrame {
ok: true,
result: None,
error: None,
namespace_mismatch: false,
config_mismatch: false,
served_config_id: Some("expected-config".to_string()),
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
error_detail: None,
metrics: None,
request_id: Some(42),
};
let payload = serde_json::to_vec(&resp).expect("encode response");
write_frame(&mut stream, &payload)
.await
.expect("write response");
}
});
let speaks = socket_speaks_khived_protocol(&sock_path, "expected-config").await;
assert!(
!speaks,
"an otherwise-matching response carrying an echoed request_id must not be treated \
as a probe acknowledgement"
);
tokio::time::timeout(std::time::Duration::from_secs(2), accept_task)
.await
.expect("fake listener accept task timed out")
.expect("fake listener accept task panicked");
}
#[derive(Clone)]
struct DetailedDispatch {
calls: Arc<std::sync::atomic::AtomicUsize>,
detail: serde_json::Value,
}
#[async_trait]
impl DaemonDispatch for DetailedDispatch {
fn plan(&self, ops: &str) -> String {
khive_request::plan_request(ops, &Default::default()).to_string()
}
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> {
panic!("the daemon must use the detailed dispatch seam");
}
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.calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Err(DaemonDispatchError::new(
"audit failed",
Some(self.detail.clone()),
))
}
async fn warm_all(&self) {}
fn namespace(&self) -> &str {
"local"
}
fn config_id(&self) -> &str {
"disposition-test"
}
}
#[tokio::test]
async fn disposition_detail_survives_daemon_framing_and_legacy_v4_decoder() {
#[allow(dead_code)]
#[derive(serde::Deserialize)]
struct LegacyV4Response {
ok: bool,
result: Option<String>,
error: Option<String>,
namespace_mismatch: bool,
#[serde(default)]
config_mismatch: bool,
#[serde(default)]
served_config_id: Option<String>,
#[serde(default)]
version_mismatch: bool,
#[serde(default)]
daemon_protocol_version: u32,
#[serde(default)]
metrics: Option<MetricsSnapshot>,
#[serde(default)]
request_id: Option<u64>,
}
let detail = serde_json::json!({
"kind": "obligation",
"code": "store_failure",
"message": "audit failed",
"domain_disposition": "committed",
"domain_result": { "id": "persisted-row" },
});
let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let response = round_trip(
DetailedDispatch {
calls: Arc::clone(&calls),
detail: detail.clone(),
},
&base_request_frame("disposition-test"),
)
.await;
assert_eq!(response.error_detail.as_ref(), Some(&detail));
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 1);
let encoded = serde_json::to_vec(&response).expect("serialize detailed response");
let legacy: LegacyV4Response = serde_json::from_slice(&encoded).expect("legacy v4 decode");
assert!(!legacy.ok);
assert_eq!(legacy.error.as_deref(), Some("audit failed"));
assert_eq!(legacy.daemon_protocol_version, PROTOCOL_VERSION);
}
#[tokio::test]
async fn disposition_legacy_dispatch_error_is_unknown_and_success_has_no_detail() {
let calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "disposition-test".to_string(),
dispatch_calls: Arc::clone(&calls),
pool: None,
dispatch_err: Some("legacy failure".to_string()),
};
let request = base_request_frame("disposition-test");
let failure = round_trip(dispatcher.clone(), &request).await;
assert_eq!(
failure.error_detail.as_ref().unwrap()["domain_disposition"],
"unknown"
);
assert_eq!(failure.error.as_deref(), Some("legacy failure"));
let success = round_trip(
MockDispatch {
dispatch_err: None,
..dispatcher
},
&request,
)
.await;
assert!(success.ok);
assert!(serde_json::to_value(success)
.unwrap()
.get("error_detail")
.is_none());
assert_eq!(calls.load(std::sync::atomic::Ordering::SeqCst), 2);
}
#[test]
fn disposition_new_decoder_accepts_legacy_v4_error_without_detail() {
let response: DaemonResponseFrame = serde_json::from_str(
r#"{
"ok":false,"result":null,"error":"legacy failure",
"namespace_mismatch":false,"config_mismatch":false,
"served_config_id":"cfg","version_mismatch":false,
"daemon_protocol_version":4,"request_id":null
}"#,
)
.expect("decode legacy v4 error frame");
assert!(response.error_detail.is_none());
assert_eq!(response.error.as_deref(), Some("legacy failure"));
}
#[test]
fn disposition_normalization_omits_unconfirmed_domain_results() {
for disposition in ["not_committed", "unknown", "unrecognized"] {
let error = DaemonDispatchError::new(
"failure",
Some(serde_json::json!({
"message": "failure",
"domain_disposition": disposition,
"domain_result": { "id": "unconfirmed" },
})),
);
assert!(error.error_detail.get("domain_result").is_none());
assert_eq!(
error.error_detail["domain_disposition"],
if disposition == "not_committed" {
"not_committed"
} else {
"unknown"
}
);
}
}
#[test]
fn disposition_normalization_iteratively_discards_deep_owned_values() {
for disposition in ["committed", "not_committed", "unknown"] {
let mut value = serde_json::Value::Null;
for _ in 0..4096 {
value = serde_json::Value::Array(vec![value]);
}
let fields = serde_json::Map::from_iter([
("domain_disposition".into(), serde_json::json!(disposition)),
("domain_result".into(), value),
]);
let error =
DaemonDispatchError::new("failure", Some(serde_json::Value::Object(fields)));
assert!(error.error_detail.get("domain_result").is_none());
assert_eq!(error.error_detail["domain_disposition"], disposition);
if disposition == "committed" {
assert_eq!(error.error_detail["code"], "result_too_deep");
}
serde_json::to_vec(&error.error_detail).expect("bounded error detail serializes");
}
let mut value = serde_json::Value::Null;
for _ in 0..4096 {
value = serde_json::Value::Array(vec![value]);
}
let error = DaemonDispatchError::new("failure", Some(value));
assert_eq!(error.error_detail["code"], "error_detail_too_deep");
assert_eq!(error.error_detail["domain_disposition"], "unknown");
assert!(error.error_detail.get("data").is_none());
}
#[tokio::test]
async fn daemon_peer_disconnect_signals_request_read_cancellation() {
let started = Arc::new(tokio::sync::Notify::new());
let cancellation_observed = Arc::new(std::sync::atomic::AtomicBool::new(false));
let dispatcher = CancellationAwareDispatch {
started: Arc::clone(&started),
cancellation_observed: Arc::clone(&cancellation_observed),
count_sql: None,
};
let (mut client, server) = UnixStream::pair().expect("unix stream pair");
let request = base_request_frame("disconnect-test");
let payload = serde_json::to_vec(&request).expect("encode request frame");
let handler = tokio::spawn(async move { handle_conn(server, dispatcher).await });
write_frame(&mut client, &payload)
.await
.expect("write request frame");
started.notified().await;
drop(client);
tokio::time::timeout(std::time::Duration::from_millis(500), handler)
.await
.expect("daemon handler ignored peer disconnect")
.expect("daemon handler panicked");
assert!(
cancellation_observed.load(std::sync::atomic::Ordering::SeqCst),
"peer loss did not reach the request-scoped read cancellation signal"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn daemon_disconnect_interrupts_pooled_stats_count() {
use khive_storage::{SqlAccess, SqlStatement, SqlValue};
let dir = tempfile::tempdir().unwrap();
let pool = Arc::new(
ConnectionPool::new(khive_db::PoolConfig {
path: Some(dir.path().join("disconnect-count.db")),
max_readers: 1,
..Default::default()
})
.unwrap(),
);
pool.writer()
.unwrap()
.conn()
.execute_batch(
"CREATE TABLE count_fixture(n INTEGER PRIMARY KEY); \
WITH RECURSIVE n(x) AS (SELECT 1 UNION ALL SELECT x+1 FROM n WHERE x<1000) \
INSERT INTO count_fixture SELECT x FROM n; \
CREATE VIEW events AS SELECT 'local' AS namespace, 'knowledge.learn' AS verb \
FROM count_fixture a CROSS JOIN count_fixture b CROSS JOIN count_fixture c;",
)
.unwrap();
let sql = Arc::new(khive_db::SqlBridge::new(Arc::clone(&pool), true));
let cancellation_observed = Arc::new(std::sync::atomic::AtomicBool::new(false));
let dispatcher = CancellationAwareDispatch {
started: Arc::new(tokio::sync::Notify::new()),
cancellation_observed: Arc::clone(&cancellation_observed),
count_sql: Some(sql.clone()),
};
let (mut client, server) = UnixStream::pair().unwrap();
let request = base_request_frame("disconnect-test");
let progress = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let handler = tokio::spawn(khive_db::scope_test_read_progress(
Arc::clone(&progress),
async move { handle_conn(server, dispatcher).await },
));
write_frame(&mut client, &serde_json::to_vec(&request).unwrap())
.await
.unwrap();
tokio::time::timeout(std::time::Duration::from_secs(2), async {
while progress.load(std::sync::atomic::Ordering::SeqCst) == 0 {
assert!(
!handler.is_finished(),
"COUNT returned before its first SQLite progress callback"
);
tokio::task::yield_now().await;
}
})
.await
.unwrap();
assert!(
!handler.is_finished(),
"COUNT must be outstanding at disconnect"
);
let started = std::time::Instant::now();
let grace = khive_db::sqlite_interrupt_grace_from_env();
drop(client);
tokio::time::timeout(grace, handler)
.await
.expect("disconnected COUNT did not settle within interrupt grace")
.unwrap();
assert!(cancellation_observed.load(std::sync::atomic::Ordering::SeqCst));
let snapshot = pool.reader_acquisition_snapshot();
assert_eq!(snapshot.active_pooled_checkouts, 0);
assert_eq!(snapshot.available_reader_admission_slots, 1);
eprintln!(
"daemon_stats_count_disconnect_ms={} grace_ms={}",
started.elapsed().as_secs_f64() * 1000.0,
grace.as_millis()
);
let count = sql
.reader()
.await
.unwrap()
.query_scalar(SqlStatement {
sql: "SELECT COUNT(*) FROM count_fixture".into(),
params: vec![],
label: None,
})
.await
.unwrap();
assert!(matches!(count, Some(SqlValue::Integer(1000))));
}
#[tokio::test]
async fn protocol_v3_frame_is_rejected_before_process_ref_dispatch() {
const {
assert!(
PROTOCOL_VERSION >= 4,
"process_ref requires protocol v4 or later"
)
};
let dispatch_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-v4".to_string(),
dispatch_calls: Arc::clone(&dispatch_calls),
pool: None,
dispatch_err: None,
};
let mut request = base_request_frame("cfg-v4");
request.protocol_version = 3;
request.process_ref = Some("worker/legacy-rollout".to_string());
let response = round_trip(dispatcher, &request).await;
assert!(!response.ok);
assert!(
!response.version_mismatch,
"a client below this protocol is answered in the implicit shape its bridge re-execs on"
);
assert_eq!(
response.error_detail.as_ref().unwrap()["code"],
"version_mismatch"
);
assert_eq!(
response.error_detail.as_ref().unwrap()["domain_disposition"],
"unknown"
);
assert_eq!(response.daemon_protocol_version, PROTOCOL_VERSION);
assert_eq!(
dispatch_calls.load(std::sync::atomic::Ordering::SeqCst),
0,
"a v3 frame must be rejected before a provenance-bearing mutation dispatches"
);
let error = response.error.expect("mismatch explains both versions");
assert!(
error.contains("client=3") && error.contains(&format!("daemon={PROTOCOL_VERSION}")),
"mismatch must identify the exact rollout boundary; got {error:?}"
);
}
#[tokio::test]
async fn protocol_v7_frame_is_rejected_before_compatible_superset_dispatch() {
let dispatch_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let client_id = config_id("p", "");
let daemon_id = config_id("p", "m");
assert!(super::config_ids_compatible(&client_id, &daemon_id));
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: daemon_id,
dispatch_calls: Arc::clone(&dispatch_calls),
pool: None,
dispatch_err: None,
};
let mut request = base_request_frame(&client_id);
request.protocol_version = 7;
let response = round_trip(dispatcher, &request).await;
assert!(!response.ok);
assert!(!response.version_mismatch);
assert_eq!(
response.error_detail.as_ref().unwrap()["code"],
"version_mismatch"
);
assert_eq!(response.daemon_protocol_version, PROTOCOL_VERSION);
assert_eq!(dispatch_calls.load(std::sync::atomic::Ordering::SeqCst), 0);
}
#[tokio::test]
async fn newer_client_frame_is_refused_with_the_explicit_flag() {
let dispatch_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-v4".to_string(),
dispatch_calls: Arc::clone(&dispatch_calls),
pool: None,
dispatch_err: None,
};
let mut request = base_request_frame("cfg-v4");
request.protocol_version = PROTOCOL_VERSION + 1;
let response = round_trip(dispatcher, &request).await;
assert!(!response.ok);
assert!(
response.version_mismatch,
"a client above this protocol keeps the explicit flag"
);
assert_eq!(response.daemon_protocol_version, PROTOCOL_VERSION);
assert_eq!(
response.error_detail.as_ref().unwrap()["code"],
"version_mismatch"
);
assert_eq!(
dispatch_calls.load(std::sync::atomic::Ordering::SeqCst),
0,
"a newer client's frame must not dispatch"
);
}
#[tokio::test]
async fn metrics_only_frame_returns_snapshot_without_dispatching() {
let dispatch_calls = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-a".to_string(),
dispatch_calls: Arc::clone(&dispatch_calls),
pool: None,
dispatch_err: None,
};
let mut metrics_req = base_request_frame("cfg-a");
metrics_req.metrics_only = true;
let metrics_resp = round_trip(dispatcher.clone(), &metrics_req).await;
assert!(metrics_resp.ok, "metrics_only response must be ok=true");
assert!(
metrics_resp.metrics.is_some(),
"metrics_only=true must return Some(snapshot)"
);
assert_eq!(
dispatch_calls.load(std::sync::atomic::Ordering::SeqCst),
0,
"metrics_only must never reach the ops-dispatch path"
);
let mut mismatched_req = base_request_frame("some-other-config");
mismatched_req.metrics_only = true;
let mismatched_resp = round_trip(dispatcher.clone(), &mismatched_req).await;
assert!(mismatched_resp.ok);
assert!(mismatched_resp.metrics.is_some());
assert!(!mismatched_resp.config_mismatch);
assert_eq!(
dispatch_calls.load(std::sync::atomic::Ordering::SeqCst),
0,
"a mismatched-config metrics_only request must still skip dispatch"
);
let normal_req = base_request_frame("cfg-a");
let normal_resp = round_trip(dispatcher, &normal_req).await;
assert!(normal_resp.ok);
assert!(normal_resp.metrics.is_none());
assert_eq!(dispatch_calls.load(std::sync::atomic::Ordering::SeqCst), 1);
}
#[test]
fn shared_writable_socket_dirs_are_refused_without_being_modified() {
let dir = tempfile::tempdir().expect("tempdir");
for (name, mode) in [("open", 0o777u32), ("sticky-tmp", 0o1777u32)] {
let shared = dir.path().join(name);
std::fs::create_dir(&shared).expect("create");
std::fs::set_permissions(&shared, std::fs::Permissions::from_mode(mode))
.expect("chmod");
let err = ensure_socket_dir_is_trusted(&shared)
.expect_err("group/other-writable must be refused, sticky or not");
assert!(
err.to_string().contains(&format!("{:04o}", mode & 0o7777)),
"the refusal should name the mode it saw, got: {err}"
);
let after = std::fs::metadata(&shared)
.expect("stat")
.permissions()
.mode()
& 0o7777;
assert_eq!(
after, mode,
"refusing must not re-permission a directory khive does not own"
);
}
}
#[test]
fn writable_non_sticky_ancestor_is_refused() {
let dir = tempfile::tempdir().expect("tempdir");
let open_mid = dir.path().join("open-mid");
std::fs::create_dir(&open_mid).expect("create mid");
let inner = open_mid.join("private");
std::fs::create_dir(&inner).expect("create inner");
std::fs::set_permissions(&inner, std::fs::Permissions::from_mode(0o700)).expect("chmod");
std::fs::set_permissions(&open_mid, std::fs::Permissions::from_mode(0o777))
.expect("chmod mid");
let err = ensure_socket_dir_is_trusted(&inner)
.expect_err("a 0777 non-sticky ancestor must be refused");
assert!(
err.to_string().contains("ancestor"),
"the refusal should say it was an ancestor that failed, got: {err}"
);
assert!(
err.to_string().contains("open-mid"),
"the refusal should name the failing ancestor, got: {err}"
);
}
#[test]
fn symlink_component_to_untrusted_directory_is_refused() {
let dir = tempfile::tempdir().expect("tempdir");
let open = dir.path().join("open-target");
std::fs::create_dir(&open).expect("create target");
std::fs::set_permissions(&open, std::fs::Permissions::from_mode(0o777)).expect("chmod");
let link = dir.path().join("link");
std::os::unix::fs::symlink(&open, &link).expect("symlink");
let euid = unsafe { libc::geteuid() } as u32;
let err = ensure_socket_path_is_swap_resistant(&link, euid)
.expect_err("a link into a 0777 non-sticky directory must be refused");
assert!(
err.to_string().contains("open-target"),
"the refusal should name the untrusted target directory, got: {err}"
);
}
include!("daemon/socket_path_tests.rs");
#[test]
fn trusted_socket_dirs_are_accepted_unmodified() {
let dir = tempfile::tempdir().expect("tempdir");
for (name, mode) in [("private", 0o700), ("listable", 0o755)] {
let d = dir.path().join(name);
std::fs::create_dir(&d).expect("create");
std::fs::set_permissions(&d, std::fs::Permissions::from_mode(mode)).expect("chmod");
ensure_socket_dir_is_trusted(&d)
.unwrap_or_else(|e| panic!("mode {mode:04o} must be accepted, got: {e}"));
let after = std::fs::metadata(&d).expect("stat").permissions().mode() & 0o7777;
assert_eq!(
after, mode,
"acceptance must not re-permission the directory either"
);
}
}
#[test]
fn pid_directory_owned_by_another_uid_is_refused() {
let dir = tempfile::Builder::new()
.prefix("khive-pid-owner-")
.tempdir()
.expect("tempdir");
let parent = dir.path().join("private");
std::fs::create_dir(&parent).expect("create private directory");
std::fs::set_permissions(&parent, std::fs::Permissions::from_mode(0o700))
.expect("set private mode");
let daemon_euid = (unsafe { libc::geteuid() } as u32).wrapping_add(1);
let error = ensure_rendezvous_dir_is_trusted(
&parent,
RendezvousPathRole::PidFile,
daemon_euid,
false,
)
.expect_err("a PID parent owned by another uid must be refused");
let message = format!("{error:#}");
assert!(
message.contains("KHIVE_PID"),
"wrong variable in refusal: {message}"
);
assert!(
message.contains("PID-file directory") && message.contains("owned by uid"),
"refusal must identify foreign ownership of the PID parent: {message}"
);
}
#[tokio::test]
async fn metrics_snapshot_wal_pages_reflects_recent_write() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("metrics_wal_test.db");
let pool = Arc::new(
ConnectionPool::new(khive_db::PoolConfig {
path: Some(path),
..khive_db::PoolConfig::for_test()
})
.expect("pool open"),
);
{
let writer = pool.try_writer().expect("writer");
writer
.conn()
.execute_batch(
"CREATE TABLE t (x INTEGER); \
INSERT INTO t VALUES (1); \
INSERT INTO t VALUES (2);",
)
.expect("seed writes");
}
let dedicated_conn = pool
.open_standalone_writer()
.expect("open dedicated checkpoint connection");
khive_db::checkpoint_once(
&pool,
&dedicated_conn,
&CheckpointConfig::default(),
&mut khive_db::checkpoint::TruncateState::default(),
)
.expect("checkpoint_once must observe on a healthy dedicated connection");
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-wal".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: Some(pool),
dispatch_err: None,
};
let snapshot = build_metrics_snapshot(&dispatcher);
assert!(
snapshot.wal_pages.is_some(),
"wal_pages must be observed after a real checkpoint tick, got {snapshot:?}"
);
assert_eq!(snapshot.wal_log_frames, snapshot.wal_pages);
assert!(snapshot.wal_checkpointed_frames.is_some());
assert!(snapshot.wal_pending_frames.is_some());
assert!(snapshot.wal_physical_bytes.is_some());
assert!(snapshot.wal_observed_at_unix_ms.is_some());
assert_eq!(snapshot.wal_checkpoint_stores.len(), 1);
assert_eq!(snapshot.wal_checkpoint_stores[0].store_id, "main");
assert_eq!(snapshot.wal_checkpoint_stores[0].timing.ticks, 1);
assert_eq!(
snapshot.wal_checkpoint_consecutive_skips, 0,
"an observed (non-skipped) tick must report zero consecutive skips, got {snapshot:?}"
);
}
#[derive(Clone)]
struct CheckpointMetricsDispatch {
main: Option<Arc<ConnectionPool>>,
secondaries: Vec<Arc<ConnectionPool>>,
}
#[async_trait]
impl DaemonDispatch for CheckpointMetricsDispatch {
fn plan(&self, _ops: &str) -> String {
panic!("metrics must not plan")
}
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> {
panic!("metrics must not dispatch")
}
async fn warm_all(&self) {}
fn namespace(&self) -> &str {
"local"
}
fn config_id(&self) -> &str {
"checkpoint-metrics"
}
fn pool_for_checkpoint(&self) -> Option<Arc<ConnectionPool>> {
self.main.clone()
}
fn secondary_pools_for_checkpoint(&self) -> Vec<Arc<ConnectionPool>> {
self.secondaries.clone()
}
}
#[tokio::test]
#[serial(checkpoint_skip_metrics)]
async fn metrics_checkpoint_timing_keeps_stores_separate_and_scrapes_read_only() {
let dir = tempfile::tempdir().unwrap();
let mut pools = Vec::new();
for label in ["primary", "secondary"] {
let directory = dir.path().join(label);
std::fs::create_dir(&directory).unwrap();
let pool = Arc::new(
ConnectionPool::new(khive_db::PoolConfig {
path: Some(directory.join("same.db")),
..khive_db::PoolConfig::for_test()
})
.unwrap(),
);
pool.try_writer()
.unwrap()
.conn()
.execute_batch("CREATE TABLE t (x INTEGER); INSERT INTO t VALUES (1);")
.unwrap();
pools.push(pool);
}
let dispatcher = CheckpointMetricsDispatch {
main: Some(Arc::clone(&pools[0])),
secondaries: vec![Arc::clone(&pools[1])],
};
let before = build_metrics_snapshot(&dispatcher);
assert_eq!(before.wal_checkpoint_stores.len(), 2);
assert!(before
.wal_checkpoint_stores
.iter()
.all(|store| store.timing.ticks == 0));
for (index, pool) in pools.iter().enumerate() {
let conn = pool.open_standalone_writer().unwrap();
for _ in 0..=index {
khive_db::checkpoint_once(
pool,
&conn,
&CheckpointConfig {
truncate_high_water_pages: u64::MAX,
..CheckpointConfig::default()
},
&mut khive_db::checkpoint::TruncateState::default(),
)
.unwrap();
}
}
let mut request = base_request_frame("checkpoint-metrics");
request.metrics_only = true;
let snapshot = round_trip(dispatcher.clone(), &request)
.await
.metrics
.unwrap();
let stores = &snapshot.wal_checkpoint_stores;
assert_eq!(stores.len(), 2);
assert_eq!(stores[0].store_id, "main");
assert_eq!(stores[0].role, "main");
assert_eq!(stores[1].store_id, "secondary:0");
assert_eq!(stores[1].role, "secondary");
for store in stores {
assert_eq!(
store.database.as_deref(),
Some("same.db"),
"wire label must omit directories"
);
assert!(store.timing.elapsed_us_max <= store.timing.elapsed_us_sum);
}
assert_eq!(stores[0].timing.ticks, 1);
assert_eq!(stores[1].timing.ticks, 2);
let again = build_metrics_snapshot(&dispatcher);
assert_eq!(
again.wal_checkpoint_stores, *stores,
"scraping must not checkpoint"
);
let secondary_only = build_metrics_snapshot(&CheckpointMetricsDispatch {
main: None,
secondaries: vec![Arc::clone(&pools[1])],
});
assert_eq!(secondary_only.wal_checkpoint_stores.len(), 1);
assert_eq!(
secondary_only.wal_checkpoint_stores[0].store_id,
"secondary:0"
);
assert_eq!(secondary_only.wal_checkpoint_stores[0].timing.ticks, 2);
assert!(build_metrics_snapshot(&CheckpointMetricsDispatch {
main: None,
secondaries: vec![]
})
.wal_checkpoint_stores
.is_empty());
}
#[test]
fn metrics_checkpoint_timing_serde_is_additive_and_round_trips() {
let snapshot = MetricsSnapshot {
wal_checkpoint_stores: vec![CheckpointStoreMetrics {
store_id: "secondary:0".into(),
role: "secondary".into(),
database: Some("memory.db".into()),
timing: khive_db::checkpoint::CheckpointTiming {
ticks: 7,
elapsed_us_sum: 123,
elapsed_us_max: 50,
busy_ticks: 2,
error_ticks: 1,
},
}],
..MetricsSnapshot::default()
};
let wire = serde_json::to_value(&snapshot).unwrap();
let store = &wire["wal_checkpoint_stores"][0];
assert_eq!(
store,
&serde_json::json!({
"store_id": "secondary:0", "role": "secondary", "database": "memory.db",
"ticks": 7, "elapsed_us_sum": 123, "elapsed_us_max": 50, "busy_ticks": 2, "error_ticks": 1,
})
);
assert_eq!(
serde_json::from_value::<MetricsSnapshot>(wire.clone()).unwrap(),
snapshot
);
let mut old_wire = wire.clone();
old_wire
.as_object_mut()
.unwrap()
.remove("wal_checkpoint_stores");
let old = serde_json::from_value::<MetricsSnapshot>(old_wire).unwrap();
assert!(
old.wal_checkpoint_stores.is_empty(),
"old snapshot must default the new vector"
);
let partial = serde_json::from_value::<CheckpointStoreMetrics>(serde_json::json!({
"store_id": "main", "role": "main"
}))
.unwrap();
assert_eq!(
partial.timing,
khive_db::checkpoint::CheckpointTiming::default()
);
assert_eq!(partial.database, None);
#[derive(serde::Deserialize)]
struct LegacyMetrics {
wal_pages: Option<u64>,
open_tx_count: usize,
}
let legacy: LegacyMetrics = serde_json::from_value(wire).unwrap();
assert_eq!(legacy.wal_pages, None);
assert_eq!(legacy.open_tx_count, 0);
}
#[tokio::test]
async fn metrics_snapshot_exposes_decomposed_writer_stages() {
let dir = tempfile::tempdir().expect("tempdir");
let path = dir.path().join("metrics_writer_stage_test.db");
let pool = Arc::new(
ConnectionPool::new(khive_db::PoolConfig {
path: Some(path),
..khive_db::PoolConfig::for_test()
})
.expect("pool open"),
);
{
let writer = pool.try_writer().unwrap();
writer
.conn()
.execute_batch("CREATE TABLE t (id INTEGER PRIMARY KEY)")
.unwrap();
}
let handle = pool
.writer_task_handle()
.unwrap()
.expect("file-backed default writer task");
handle
.send(|conn| {
std::thread::sleep(std::time::Duration::from_millis(30));
conn.execute("INSERT INTO t VALUES (1)", [])
.map_err(|error| khive_storage::error::StorageError::Pool {
operation: "metrics_writer_stage_test".into(),
message: error.to_string(),
})
})
.await
.unwrap();
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-writer-stages".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: Some(pool),
dispatch_err: None,
};
let snapshot = build_metrics_snapshot(&dispatcher);
assert!(snapshot.write_last_queue_wait_micros.is_some());
assert!(snapshot.write_last_transaction_acquire_micros.is_some());
assert!(snapshot.write_last_commit_micros.is_some());
assert!(
snapshot.write_last_body_micros >= Some(25_000),
"synthetic delay must be attributed to the body: {snapshot:?}"
);
assert!(snapshot.write_last_total_micros >= snapshot.write_last_body_micros);
assert!(snapshot.write_last_observed_at_unix_ms.is_some());
}
#[test]
#[serial(tx_registry)]
fn metrics_snapshot_reflects_open_transaction_registry() {
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-tx".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: None,
dispatch_err: None,
};
let departing_handle = khive_storage::tx_registry::register(Some(
"daemon_metrics_snapshot_departing_test_tx".to_string(),
));
let before = build_metrics_snapshot(&dispatcher).open_tx_count;
assert!(before >= 1);
let handle = khive_storage::tx_registry::register(Some(
"daemon_metrics_snapshot_owned_test_tx".to_string(),
));
drop(departing_handle);
let during = build_metrics_snapshot(&dispatcher);
assert!(
during.open_tx_count >= 1,
"open_tx_count must reflect the live owned transaction despite registry churn: \
churn_baseline={before} during={}",
during.open_tx_count
);
assert!(
during.oldest_pinned_tx_micros.is_some(),
"oldest_pinned_tx_micros must be Some while a transaction is open"
);
drop(handle);
assert!(
!khive_storage::tx_registry::snapshot()
.iter()
.any(|(_, label)| label.as_deref()
== Some("daemon_metrics_snapshot_owned_test_tx")),
"the owned registry entry must disappear when its handle is dropped"
);
}
#[tokio::test]
async fn metrics_snapshot_write_queue_depth_flag_gated() {
let dir = tempfile::tempdir().expect("tempdir");
let enabled_pool = Arc::new(
ConnectionPool::new(khive_db::PoolConfig {
path: Some(dir.path().join("wq_enabled.db")),
write_queue_enabled: Some(true),
..khive_db::PoolConfig::for_test()
})
.expect("pool open"),
);
let enabled_dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-wq-on".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: Some(enabled_pool),
dispatch_err: None,
};
let snapshot_on = build_metrics_snapshot(&enabled_dispatcher);
assert!(
snapshot_on.write_queue_depth.is_some(),
"write_queue_depth must be Some when write_queue_enabled=true, got {snapshot_on:?}"
);
assert!(snapshot_on.write_queue_capacity.is_some());
let disabled_pool = Arc::new(
ConnectionPool::new(khive_db::PoolConfig {
path: Some(dir.path().join("wq_disabled.db")),
write_queue_enabled: Some(false),
..khive_db::PoolConfig::for_test()
})
.expect("pool open"),
);
let disabled_dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-wq-off".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: Some(disabled_pool),
dispatch_err: None,
};
let snapshot_off = build_metrics_snapshot(&disabled_dispatcher);
assert!(
snapshot_off.write_queue_depth.is_none(),
"write_queue_depth must be None when write_queue_enabled=false, got {snapshot_off:?}"
);
assert!(snapshot_off.write_queue_capacity.is_none());
let no_pool_dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-no-pool".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: None,
dispatch_err: None,
};
let snapshot_no_pool = build_metrics_snapshot(&no_pool_dispatcher);
assert!(snapshot_no_pool.write_queue_depth.is_none());
assert!(snapshot_no_pool.write_queue_capacity.is_none());
}
#[test]
fn frame_serde_defaults_additive_fields_when_absent() {
let req_json = serde_json::json!({
"ops": "",
"presentation": null,
"presentation_per_op": null,
"namespace": "local",
"actor_id": null,
"visible_namespaces": [],
"config_id": "cfg",
"protocol_version": PROTOCOL_VERSION,
"probe_only": false,
"format": null,
"format_per_op": null,
"from_wire": false
});
let frame: DaemonRequestFrame =
serde_json::from_value(req_json).expect("decode a metrics_only-absent request frame");
assert!(
!frame.metrics_only,
"metrics_only must default to false when absent from the wire payload"
);
assert_eq!(
frame.request_id, None,
"request_id must default to None when absent from the wire payload (khive#948)"
);
assert_eq!(
frame.process_ref, None,
"process_ref must default to None when absent from the wire payload (khive#1428)"
);
let encoded_frame = serde_json::to_value(&frame).expect("encode request frame");
assert!(
encoded_frame.get("process_ref").is_none(),
"absent provenance must not change the serialized request wire shape"
);
let resp_json = serde_json::json!({
"ok": true,
"result": null,
"error": null,
"namespace_mismatch": false,
"config_mismatch": false,
"served_config_id": "cfg",
"version_mismatch": false,
"daemon_protocol_version": PROTOCOL_VERSION
});
let resp: DaemonResponseFrame =
serde_json::from_value(resp_json).expect("decode a metrics-absent response frame");
assert!(
resp.metrics.is_none(),
"metrics must default to None when absent from the wire payload"
);
assert_eq!(
resp.request_id, None,
"request_id must default to None when absent from the wire payload (khive#948)"
);
}
#[tokio::test]
async fn request_id_echoed_on_success_and_error_arms() {
let dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-a".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: None,
dispatch_err: None,
};
let mut ok_req = base_request_frame("cfg-a");
ok_req.request_id = Some(42);
let ok_resp = round_trip(dispatcher, &ok_req).await;
assert!(ok_resp.ok, "expected successful dispatch: {ok_resp:?}");
assert_eq!(
ok_resp.request_id,
Some(42),
"request_id must be echoed back on a successful dispatch response"
);
let mismatched_dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-a".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: None,
dispatch_err: None,
};
let mut mismatch_req = base_request_frame("cfg-WRONG");
mismatch_req.request_id = Some(99);
let mismatch_resp = round_trip(mismatched_dispatcher, &mismatch_req).await;
assert!(mismatch_resp.config_mismatch);
assert_eq!(
mismatch_resp.request_id,
Some(99),
"request_id must be echoed on the config_mismatch rejection arm too"
);
let erroring_dispatcher = MockDispatch {
namespace: "local".to_string(),
config_id: "cfg-a".to_string(),
dispatch_calls: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
pool: None,
dispatch_err: Some("simulated dispatch error".to_string()),
};
let mut err_req = base_request_frame("cfg-a");
err_req.request_id = Some(7);
let err_resp = round_trip(erroring_dispatcher, &err_req).await;
assert!(!err_resp.ok, "expected a dispatch error: {err_resp:?}");
assert_eq!(
err_resp.request_id,
Some(7),
"request_id must be echoed on the real ops-dispatch error arm"
);
}
#[test]
fn shutdown_cleanup_removes_paths_it_still_owns() {
let dir = tempfile::tempdir().expect("tempdir");
let sock = dir.path().join("khived.sock");
let pid_file = dir.path().join("khived.pid");
let _listener = std::os::unix::net::UnixListener::bind(&sock).expect("bind socket");
std::fs::write(&pid_file, std::process::id().to_string()).expect("write pid file");
let identity = socket_identity(&sock);
assert!(
identity.is_some(),
"must read identity of a freshly bound socket"
);
let cleaned = shutdown_cleanup_if_owned(&sock, &pid_file, identity);
assert!(
cleaned,
"cleanup must proceed when PID and socket still match"
);
assert!(!sock.exists(), "owned socket must be removed");
assert!(!pid_file.exists(), "owned pid file must be removed");
}
#[test]
fn shutdown_cleanup_skips_when_pid_file_names_a_different_process() {
let dir = tempfile::tempdir().expect("tempdir");
let sock = dir.path().join("khived.sock");
let pid_file = dir.path().join("khived.pid");
let _listener = std::os::unix::net::UnixListener::bind(&sock).expect("bind socket");
let identity = socket_identity(&sock);
std::fs::write(&pid_file, "1").expect("write foreign pid file");
let cleaned = shutdown_cleanup_if_owned(&sock, &pid_file, identity);
assert!(
!cleaned,
"cleanup must be skipped when the PID file no longer names this process"
);
assert!(sock.exists(), "replacement daemon's socket must survive");
assert!(
pid_file.exists(),
"replacement daemon's pid file must survive"
);
}
#[test]
fn shutdown_cleanup_skips_when_socket_was_rebound_by_a_replacement() {
let dir = tempfile::tempdir().expect("tempdir");
let sock = dir.path().join("khived.sock");
let original_sock = dir.path().join("original.sock");
let pid_file = dir.path().join("khived.pid");
let _original_listener =
std::os::unix::net::UnixListener::bind(&original_sock).expect("bind original socket");
let _replacement_listener =
std::os::unix::net::UnixListener::bind(&sock).expect("bind replacement socket");
let original_identity = socket_identity(&original_sock);
let replacement_identity = socket_identity(&sock);
assert!(
original_identity.is_some(),
"must read identity of the original socket"
);
assert!(
replacement_identity.is_some(),
"must read identity of the replacement socket"
);
assert!(
original_identity != replacement_identity,
"two concurrently bound sockets must have distinct identities"
);
std::fs::write(&pid_file, std::process::id().to_string())
.expect("write pid file matching this process");
let cleaned = shutdown_cleanup_if_owned(&sock, &pid_file, original_identity);
assert!(
!cleaned,
"cleanup must be skipped when the socket at this path is a different \
inode than the one this daemon originally bound"
);
assert!(sock.exists(), "replacement daemon's socket must survive");
assert!(
pid_file.exists(),
"replacement daemon's pid file must survive"
);
}
#[test]
fn shutdown_cleanup_preserves_atomically_renamed_successor() {
let dir = tempfile::tempdir().expect("tempdir");
let sock = dir.path().join("khived.sock");
let staged_sock = dir.path().join("next.sock");
let pid_file = dir.path().join("khived.pid");
let _original_listener =
std::os::unix::net::UnixListener::bind(&sock).expect("bind original socket");
let successor =
std::os::unix::net::UnixListener::bind(&staged_sock).expect("bind staged successor");
let original_identity = socket_identity(&sock).expect("original socket identity");
let successor_identity = socket_identity(&staged_sock).expect("successor socket identity");
assert!(original_identity != successor_identity);
let original_pid = std::process::id().to_string();
std::fs::write(&pid_file, &original_pid).expect("write original PID");
std::fs::rename(&staged_sock, &sock).expect("publish successor over original socket");
assert!(!staged_sock.exists());
assert!(socket_identity(&sock) == Some(successor_identity));
assert!(!shutdown_cleanup_if_owned(
&sock,
&pid_file,
Some(original_identity)
));
assert!(socket_identity(&sock) == Some(successor_identity));
assert_eq!(
std::fs::read_to_string(&pid_file).expect("PID must survive stale cleanup"),
original_pid
);
successor
.set_nonblocking(true)
.expect("bound successor must support nonblocking accept");
let _client = std::os::unix::net::UnixStream::connect(&sock)
.expect("published successor must remain reachable");
let _accepted = successor
.accept()
.expect("successor must receive connection");
}
#[test]
fn isolated_daemon_locks_use_private_fixture_paths() {
if crate::test_process::run_in_child() {
return;
}
let home = PathBuf::from(std::env::var_os("HOME").expect("child HOME"));
for path in [lock_path(), recoverer_lock_path()] {
assert_eq!(
path.parent(),
home.parent(),
"runtime daemon locks must use private fixture paths outside HOME"
);
}
let _boot = acquire_daemon_boot_guard().expect("private boot lock");
let _recoverer = try_acquire_recoverer_lock_until(
std::time::Instant::now() + std::time::Duration::from_secs(1),
)
.expect("private recoverer lock")
.expect("private recoverer lock must be available");
assert!(lock_path().is_file());
assert!(recoverer_lock_path().is_file());
assert!(
std::fs::read_dir(home).unwrap().next().is_none(),
"both daemon lock producers must leave the child HOME empty"
);
}
include!("daemon/store_guard_tests.rs");
#[test]
#[serial]
fn recovery_lock_serializes_two_concurrent_boot_sequences() {
if crate::test_process::run_in_child() {
return;
}
let dir = tempfile::tempdir().expect("tempdir");
let lock_file = dir.path().join("khived.recovery.lock");
std::env::set_var("KHIVE_LOCK", &lock_file);
let active = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let overlap_detected = Arc::new(std::sync::atomic::AtomicBool::new(false));
let run_one_boot =
|active: Arc<std::sync::atomic::AtomicUsize>,
overlap: Arc<std::sync::atomic::AtomicBool>| {
move || {
let _guard = acquire_recovery_lock().expect("acquire recovery lock");
if active.fetch_add(1, std::sync::atomic::Ordering::SeqCst) != 0 {
overlap.store(true, std::sync::atomic::Ordering::SeqCst);
}
std::thread::sleep(std::time::Duration::from_millis(50));
active.fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
}
};
let t1 = std::thread::spawn(run_one_boot(active.clone(), overlap_detected.clone()));
let t2 = std::thread::spawn(run_one_boot(active.clone(), overlap_detected.clone()));
t1.join().expect("boot thread 1 must not panic");
t2.join().expect("boot thread 2 must not panic");
assert!(
!overlap_detected.load(std::sync::atomic::Ordering::SeqCst),
"two concurrent boot sequences must never hold the schema-init \
critical section at the same time (#667)"
);
std::env::remove_var("KHIVE_LOCK");
}
#[test]
#[serial]
fn acquire_daemon_boot_guard_returns_guard_when_lock_available() {
if crate::test_process::run_in_child() {
return;
}
let dir = tempfile::tempdir().expect("tempdir");
let lock_file = dir.path().join("khived.recovery.lock");
std::env::set_var("KHIVE_LOCK", &lock_file);
let guard = acquire_daemon_boot_guard();
assert!(
guard.is_ok(),
"daemon boot guard must succeed when the lock file can be opened and flocked"
);
drop(guard);
std::env::remove_var("KHIVE_LOCK");
}
#[test]
#[serial]
fn acquire_daemon_boot_guard_fails_loudly_when_lock_file_cannot_be_opened() {
if crate::test_process::run_in_child() {
return;
}
let dir = tempfile::tempdir().expect("tempdir");
std::env::set_var("KHIVE_LOCK", dir.path());
let result = acquire_daemon_boot_guard();
assert!(
result.is_err(),
"daemon boot guard must fail loudly, never silently proceed unguarded, \
when the underlying recovery lock cannot be acquired"
);
std::env::remove_var("KHIVE_LOCK");
}
#[test]
fn write_pid_file_exclusive_creates_new_file_with_own_pid() {
let dir = tempfile::tempdir().expect("tempdir");
let pid_file = dir.path().join("khived.pid");
write_pid_file_exclusive(&pid_file).expect("first writer must win");
let contents = std::fs::read_to_string(&pid_file).expect("read pid file");
assert_eq!(contents, std::process::id().to_string());
}
#[test]
fn write_pid_file_exclusive_refuses_to_overwrite_an_existing_file() {
let dir = tempfile::tempdir().expect("tempdir");
let pid_file = dir.path().join("khived.pid");
std::fs::write(&pid_file, "999999").expect("seed an existing pid file");
let err = write_pid_file_exclusive(&pid_file)
.expect_err("must not silently overwrite an existing pid file");
assert_eq!(err.kind(), std::io::ErrorKind::AlreadyExists);
let contents = std::fs::read_to_string(&pid_file).expect("read pid file");
assert_eq!(
contents, "999999",
"an existing pid file must never be truncated by a losing writer"
);
}
#[test]
fn two_concurrent_writers_converge_on_exactly_one_pid_file_owner() {
let dir = tempfile::tempdir().expect("tempdir");
let pid_file = std::sync::Arc::new(dir.path().join("khived.pid"));
let barrier = std::sync::Arc::new(std::sync::Barrier::new(2));
let spawn_writer =
|pid_file: std::sync::Arc<std::path::PathBuf>,
barrier: std::sync::Arc<std::sync::Barrier>| {
std::thread::spawn(move || {
barrier.wait();
write_pid_file_exclusive(&pid_file)
})
};
let t1 = spawn_writer(pid_file.clone(), barrier.clone());
let t2 = spawn_writer(pid_file.clone(), barrier.clone());
let r1 = t1.join().expect("writer 1 must not panic");
let r2 = t2.join().expect("writer 2 must not panic");
let results = [&r1, &r2];
let ok_count = results.iter().filter(|r| r.is_ok()).count();
let already_exists_count = results
.iter()
.filter(|r| matches!(r, Err(e) if e.kind() == std::io::ErrorKind::AlreadyExists))
.count();
assert_eq!(
ok_count, 1,
"exactly one of two concurrent writers must win the pid file"
);
assert_eq!(
already_exists_count, 1,
"the other writer must observe AlreadyExists, never a silent overwrite"
);
assert!(pid_file.exists(), "the winner's pid file must exist");
let contents = std::fs::read_to_string(&*pid_file).expect("read pid file");
assert_eq!(
contents,
std::process::id().to_string(),
"the surviving pid file must contain the winner's pid — both threads \
share this process's pid, so an unexpected value would also prove a \
lost/garbled write raced through"
);
}
#[tokio::test]
async fn peer_uid_reports_the_connecting_process_uid() {
let dir = tempfile::tempdir().expect("tempdir");
let sock = dir.path().join("peer.sock");
let listener = UnixListener::bind(&sock).expect("bind");
let connect_path = sock.clone();
let client = tokio::spawn(async move { UnixStream::connect(&connect_path).await });
let (server_side, _) = listener.accept().await.expect("accept");
let client_side = client.await.expect("join").expect("connect");
let expected = unsafe { libc::geteuid() } as u32;
assert_eq!(
peer_uid(&server_side).expect("peer_uid must succeed on a live connection"),
expected,
"the uid read from the kernel for a same-process connection must be \
this process's euid"
);
assert_eq!(
peer_uid(&client_side).expect("peer_uid must succeed on the client end"),
expected
);
}
#[test]
fn only_a_foreign_uid_is_refused() {
let euid = unsafe { libc::geteuid() } as u32;
assert!(
uid_is_permitted(euid, euid),
"a connection from the daemon's own uid must be served — this is \
every seat on the host, and ADR-096 accepted exactly this shape"
);
assert!(
!uid_is_permitted(euid.wrapping_add(1), euid),
"a connection from any other uid must be refused"
);
assert!(
!uid_is_permitted(0, euid.wrapping_add(1)),
"root is not special-cased: the rule is equality with the daemon's \
euid, not a privilege comparison"
);
}
struct CapturedFields(Arc<std::sync::Mutex<Vec<String>>>);
impl tracing::Subscriber for CapturedFields {
fn enabled(&self, _: &tracing::Metadata<'_>) -> bool {
true
}
fn new_span(&self, _: &tracing::span::Attributes<'_>) -> tracing::span::Id {
tracing::span::Id::from_u64(1)
}
fn record(&self, _: &tracing::span::Id, _: &tracing::span::Record<'_>) {}
fn record_follows_from(&self, _: &tracing::span::Id, _: &tracing::span::Id) {}
fn event(&self, event: &tracing::Event<'_>) {
struct Visitor(String);
impl tracing::field::Visit for Visitor {
fn record_debug(
&mut self,
field: &tracing::field::Field,
value: &dyn std::fmt::Debug,
) {
self.0.push_str(&format!("{}={:?} ", field.name(), value));
}
}
let mut visitor = Visitor(String::new());
event.record(&mut visitor);
self.0.lock().unwrap().push(visitor.0);
}
fn enter(&self, _: &tracing::span::Id) {}
fn exit(&self, _: &tracing::span::Id) {}
}
#[cfg(unix)]
#[tokio::test]
#[serial(background_tasks)]
async fn drain_timeout_warning_names_the_outstanding_tasks() {
let lines = Arc::new(std::sync::Mutex::new(Vec::new()));
let subscriber = CapturedFields(lines.clone());
let _dispatch = tracing::dispatcher::set_default(&tracing::Dispatch::new(subscriber));
let active = std::sync::atomic::AtomicUsize::new(0);
let (stop_tx, stop_rx) = tokio::sync::broadcast::channel::<()>(1);
for name in ["test_task_alpha", "test_task_beta"] {
let mut rx = stop_rx.resubscribe();
track_named_background_task(name, async move {
let _ = rx.recv().await;
});
}
drop(stop_rx);
let drained = drain_with_timeout(&active, std::time::Duration::from_millis(150)).await;
assert!(
!drained,
"two unfinished tasks must make the drain time out"
);
let warned = lines
.lock()
.unwrap()
.iter()
.find(|line| line.contains("drain timeout reached"))
.cloned()
.expect("the drain timeout must emit its warning through the test subscriber");
assert!(
warned.contains("test_task_alpha") && warned.contains("test_task_beta"),
"the drain-timeout warning must name every outstanding task; got {warned}"
);
let _ = stop_tx.send(());
for _ in 0..100 {
if background_task_names().is_empty() {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
}
#[tokio::test]
#[serial(background_tasks)]
async fn a_named_task_drops_its_name_when_it_finishes() {
let before = background_task_count();
let (tx, rx) = tokio::sync::oneshot::channel::<()>();
track_named_background_task("test_task_finishes", async move {
let _ = rx.await;
});
assert!(
background_task_names().contains(&"test_task_finishes".to_string()),
"a live named task must be listed while the counter holds it"
);
tx.send(()).expect("still awaiting");
for _ in 0..100 {
if background_task_count() == before {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
assert_eq!(background_task_count(), before);
assert!(
!background_task_names().contains(&"test_task_finishes".to_string()),
"a finished task's name must be released, not left to accumulate"
);
}
#[tokio::test]
#[serial(background_tasks)]
async fn the_unnamed_entry_point_still_registers_and_releases() {
let before = background_task_count();
let (tx, rx) = tokio::sync::oneshot::channel::<()>();
track_background_task(async move {
let _ = rx.await;
});
assert_eq!(background_task_count(), before + 1);
assert!(
background_task_names().contains(&UNNAMED_BACKGROUND_TASK.to_string()),
"the unchanged public entry point must still register, under the placeholder name"
);
tx.send(()).expect("still awaiting");
for _ in 0..100 {
if background_task_count() == before {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
assert_eq!(background_task_count(), before);
}
include!("daemon_config_id_tests.rs");
}