use std::sync::Arc;
use serde::{Deserialize, Serialize};
#[cfg(unix)]
use tokio::net::UnixStream;
use super::spawn_named_tracked_task;
#[cfg(unix)]
use super::{write_frame, DaemonResponseFrame, PROTOCOL_VERSION};
#[cfg(unix)]
const MAX_CONNECTIONS_ENV: &str = "KHIVE_DAEMON_MAX_CONNECTIONS";
#[cfg(unix)]
const DEFAULT_MAX_CONNECTIONS: usize = 512;
#[cfg(unix)]
const RESERVED_DESCRIPTORS: u64 = 192;
#[cfg(unix)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(super) enum CapSource {
Builtin,
Configured,
DescriptorLimit,
}
#[cfg(unix)]
const BUSY_REFUSAL_WRITE_TIMEOUT: std::time::Duration = std::time::Duration::from_millis(250);
pub(super) fn parse_positive_limit(raw: Option<&str>) -> Option<u64> {
raw.and_then(|value| value.trim().parse::<u64>().ok())
.filter(|value| *value > 0)
}
fn positive_limit_from_env(name: &str) -> Option<u64> {
let raw = std::env::var(name).ok()?;
let parsed = parse_positive_limit(Some(&raw));
if parsed.is_none() {
tracing::warn!(
variable = name,
value = %raw,
"ignoring a limit that is not a positive integer; the default applies"
);
}
parsed
}
#[cfg(unix)]
fn positive_usize_from_env(name: &str) -> Option<usize> {
let value = positive_limit_from_env(name)?;
Some(usize::try_from(value).unwrap_or(usize::MAX))
}
#[cfg(unix)]
pub(super) fn soft_nofile_limit() -> Option<u64> {
let mut limit = libc::rlimit {
rlim_cur: 0,
rlim_max: 0,
};
if unsafe { libc::getrlimit(libc::RLIMIT_NOFILE, &mut limit) } != 0 {
return None;
}
#[allow(clippy::unnecessary_cast)]
let soft = limit.rlim_cur as u64;
Some(soft)
}
#[cfg(unix)]
pub(super) fn effective_connection_cap(
configured: Option<usize>,
soft_nofile: Option<u64>,
) -> (usize, CapSource) {
let (wanted, source) = match configured {
Some(value) => (value, CapSource::Configured),
None => (DEFAULT_MAX_CONNECTIONS, CapSource::Builtin),
};
let room = match soft_nofile {
Some(soft) => soft.saturating_sub(RESERVED_DESCRIPTORS),
None => u64::MAX,
};
let room = usize::try_from(room).unwrap_or(usize::MAX);
let reduced = room < wanted;
let cap = if reduced { room } else { wanted };
let source = if reduced {
CapSource::DescriptorLimit
} else {
source
};
(cap.max(1), source)
}
#[cfg(unix)]
pub(super) struct ConnectionAdmission {
limit: usize,
configured: Option<usize>,
permits: Arc<tokio::sync::Semaphore>,
refused: std::sync::atomic::AtomicU64,
}
#[cfg(unix)]
impl ConnectionAdmission {
pub(super) fn new(limit: usize) -> Self {
let limit = limit.clamp(1, tokio::sync::Semaphore::MAX_PERMITS);
Self {
limit,
configured: None,
permits: Arc::new(tokio::sync::Semaphore::new(limit)),
refused: std::sync::atomic::AtomicU64::new(0),
}
}
pub(super) fn from_limits(configured: Option<usize>, soft_nofile: Option<u64>) -> Self {
let (limit, source) = effective_connection_cap(configured, soft_nofile);
let mut admission = Self::new(limit);
match soft_nofile {
None => tracing::warn!(
cap = admission.limit,
"could not read the descriptor limit; the connection cap is not reduced to fit it"
),
Some(soft) if source == CapSource::DescriptorLimit => {
let wanted = configured.unwrap_or(DEFAULT_MAX_CONNECTIONS);
admission.configured = Some(wanted);
tracing::warn!(
configured_cap = wanted,
effective_cap = admission.limit,
soft_nofile = soft,
reserved_descriptors = RESERVED_DESCRIPTORS,
"connection cap reduced to fit the process descriptor limit"
);
}
Some(_) => {}
}
admission
}
pub(super) fn from_env() -> Self {
Self::from_limits(
positive_usize_from_env(MAX_CONNECTIONS_ENV),
soft_nofile_limit(),
)
}
pub(super) fn try_admit(&self) -> Result<tokio::sync::OwnedSemaphorePermit, u64> {
Arc::clone(&self.permits).try_acquire_owned().map_err(|_| {
self.refused
.fetch_add(1, std::sync::atomic::Ordering::SeqCst)
+ 1
})
}
pub(super) fn snapshot(&self) -> ConnectionCapSnapshot {
ConnectionCapSnapshot {
limit: self.limit,
configured: self.configured,
active: self.limit - self.permits.available_permits(),
refused: self.refused.load(std::sync::atomic::Ordering::SeqCst),
}
}
}
#[cfg(unix)]
pub(super) async fn admit_or_refuse_busy(
admission: &ConnectionAdmission,
stream: &mut UnixStream,
config_id: &str,
) -> Option<tokio::sync::OwnedSemaphorePermit> {
let refused_total = match admission.try_admit() {
Ok(permit) => return Some(permit),
Err(refused_total) => refused_total,
};
if refused_total.is_power_of_two() {
tracing::warn!(
cap = admission.limit,
refused_total,
"daemon at connection cap; refusing new connections"
);
}
let limit = admission.limit;
let refusal = DaemonResponseFrame {
ok: false,
result: None,
error: Some(format!(
"daemon is at its connection limit ({limit}); request was not admitted, retry shortly"
)),
error_detail: Some(serde_json::json!({
"kind": "runtime", "code": "daemon_busy", "limit": limit,
"domain_disposition": crate::DomainDisposition::NotCommitted.as_str(),
})),
namespace_mismatch: false,
config_mismatch: false,
served_config_id: Some(config_id.to_owned()),
version_mismatch: false,
daemon_protocol_version: PROTOCOL_VERSION,
metrics: None,
request_id: None,
};
if let Ok(payload) = serde_json::to_vec(&refusal) {
let write = write_frame(stream, &payload);
match tokio::time::timeout(BUSY_REFUSAL_WRITE_TIMEOUT, write).await {
Ok(Ok(())) => {}
Ok(Err(error)) => tracing::debug!(%error, "failed to write connection-limit refusal"),
Err(_) => tracing::debug!("connection-limit refusal write timed out"),
}
}
None
}
#[derive(Serialize, Deserialize, Debug, Clone, Default, PartialEq, Eq)]
pub struct ConnectionCapSnapshot {
pub limit: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub configured: Option<usize>,
pub active: usize,
pub refused: u64,
}
const RECALL_LEDGER_MAX_PENDING_ENV: &str = "KHIVE_RECALL_LEDGER_MAX_PENDING";
const RECALL_LEDGER_TIMEOUT_MS_ENV: &str = "KHIVE_RECALL_LEDGER_TIMEOUT_MS";
const DEFAULT_RECALL_LEDGER_MAX_PENDING: usize = 256;
const DEFAULT_RECALL_LEDGER_TIMEOUT_MS: u64 = 30_000;
pub(super) struct RecallLedgerBound {
max_pending: usize,
timeout: std::time::Duration,
permits: Arc<tokio::sync::Semaphore>,
skipped: std::sync::atomic::AtomicU64,
timed_out: std::sync::atomic::AtomicU64,
}
impl RecallLedgerBound {
pub(super) fn new(max_pending: usize, timeout: std::time::Duration) -> Self {
let max_pending = max_pending.clamp(1, tokio::sync::Semaphore::MAX_PERMITS);
Self {
max_pending,
timeout,
permits: Arc::new(tokio::sync::Semaphore::new(max_pending)),
skipped: std::sync::atomic::AtomicU64::new(0),
timed_out: std::sync::atomic::AtomicU64::new(0),
}
}
fn from_env() -> Self {
let max_pending = match positive_limit_from_env(RECALL_LEDGER_MAX_PENDING_ENV) {
Some(value) => usize::try_from(value).unwrap_or(usize::MAX),
None => DEFAULT_RECALL_LEDGER_MAX_PENDING,
};
let timeout_ms = positive_limit_from_env(RECALL_LEDGER_TIMEOUT_MS_ENV)
.unwrap_or(DEFAULT_RECALL_LEDGER_TIMEOUT_MS);
Self::new(max_pending, std::time::Duration::from_millis(timeout_ms))
}
pub(super) fn try_spawn<F>(self: &Arc<Self>, fut: F) -> Option<tokio::task::JoinHandle<()>>
where
F: std::future::Future<Output = ()> + Send + 'static,
{
let Ok(permit) = Arc::clone(&self.permits).try_acquire_owned() else {
self.skipped
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
return None;
};
let bound = Arc::clone(self);
let handle = spawn_named_tracked_task("memory_recall_serve_ledger", async move {
let _permit = permit;
if tokio::time::timeout(bound.timeout, fut).await.is_err() {
bound
.timed_out
.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
tracing::warn!(
timeout_ms = bound.timeout.as_millis() as u64,
"recall serve ledger task timed out; the ledger write was missed"
);
}
});
Some(handle)
}
pub(super) fn snapshot(&self) -> RecallLedgerSnapshot {
RecallLedgerSnapshot {
max_pending: self.max_pending,
pending: self.max_pending - self.permits.available_permits(),
timeout_ms: self.timeout.as_millis().min(u128::from(u64::MAX)) as u64,
skipped: self.skipped.load(std::sync::atomic::Ordering::SeqCst),
timed_out: self.timed_out.load(std::sync::atomic::Ordering::SeqCst),
}
}
}
fn recall_ledger() -> &'static Arc<RecallLedgerBound> {
static BOUND: std::sync::OnceLock<Arc<RecallLedgerBound>> = std::sync::OnceLock::new();
BOUND.get_or_init(|| Arc::new(RecallLedgerBound::from_env()))
}
pub fn track_recall_ledger_task<F>(fut: F)
where
F: std::future::Future<Output = ()> + Send + 'static,
{
drop(recall_ledger().try_spawn(fut));
}
pub fn recall_ledger_snapshot() -> RecallLedgerSnapshot {
recall_ledger().snapshot()
}
#[derive(Serialize, Deserialize, Debug, Clone, Default, PartialEq, Eq)]
pub struct RecallLedgerSnapshot {
pub max_pending: usize,
pub pending: usize,
pub timeout_ms: u64,
pub skipped: u64,
pub timed_out: u64,
}