use std::collections::HashMap;
use std::ffi::OsString;
use std::process::Stdio;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use car_inference::stream::StreamEvent;
use car_inference::tasks::generate::GenerateRequest;
use car_inference::{
InferenceConfig, InferenceEngine, InferenceError, InferenceResult, LocalGenerationOffload,
LocalLoadPreflight, LocalOffloadResult, LocalOffloadStream, LocalWorkerAdmission,
LocalWorkerResidency,
};
use serde::{Deserialize, Serialize};
use tokio::io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader, Lines};
use tokio::process::{Child, ChildStdin, ChildStdout, Command};
use tokio::sync::Mutex;
pub const WORKER_ENV: &str = "CAR_INFERENCE_WORKER";
#[derive(Serialize, Deserialize)]
enum WorkerRequest {
Generate {
request: Box<GenerateRequest>,
admission: LocalWorkerAdmission,
},
Stream {
request: Box<GenerateRequest>,
admission: LocalWorkerAdmission,
},
}
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
struct WorkerResidencyAck {
model_id: String,
measured_weights_bytes: u64,
retention: car_inference::backend_cache::BackendRetention,
}
#[derive(Serialize, Deserialize)]
enum WorkerResponse {
Result {
result: Box<InferenceResult>,
residency: WorkerResidencyAck,
},
StreamStarted {
residency: WorkerResidencyAck,
},
Event(Box<StreamEvent>),
StreamEnd,
Error(String),
LocalResourceBlocked {
preflight: LocalLoadPreflight,
recovery: String,
},
}
enum Exchange<T> {
Ok(T),
Reported(InferenceError),
Dead(String),
}
struct WorkerProc {
child: Child,
stdin: ChildStdin,
stdout: Lines<BufReader<ChildStdout>>,
policy_generation: u64,
state_root: Option<std::path::PathBuf>,
}
#[derive(Clone)]
struct WorkerResident {
allocation_id: String,
coordinator: Arc<car_inference::resource_policy::LocalAdmissionCoordinator>,
}
type WorkerResidentMap = HashMap<(std::path::PathBuf, String), WorkerResident>;
#[derive(Clone)]
struct WorkerResidentOwner {
root: std::path::PathBuf,
logical_model_id: String,
allocation_id: String,
coordinator: Arc<car_inference::resource_policy::LocalAdmissionCoordinator>,
}
struct WorkerProcessGuard {
worker: Option<WorkerProc>,
residents: Vec<WorkerResidentOwner>,
candidate: Option<WorkerResidentOwner>,
resident_models: Arc<std::sync::Mutex<WorkerResidentMap>>,
teardown_started: bool,
}
impl WorkerProcessGuard {
fn new(
worker: WorkerProc,
resident_models: Arc<std::sync::Mutex<WorkerResidentMap>>,
candidate: Option<(std::path::PathBuf, String, String)>,
) -> Self {
let mut residents = resident_models
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.iter()
.map(|((root, logical_model_id), resident)| WorkerResidentOwner {
root: root.clone(),
logical_model_id: logical_model_id.clone(),
allocation_id: resident.allocation_id.clone(),
coordinator: resident.coordinator.clone(),
})
.collect::<Vec<_>>();
let mut candidate_owner = None;
if let Some((root, model_id, allocation_id)) =
candidate.filter(|(_, model_id, _)| !model_id.is_empty())
{
let root = car_inference::resource_policy::normalized_state_root_key(&root);
if !residents
.iter()
.any(|resident| resident.root == root && resident.logical_model_id == model_id)
{
if let Some(coordinator) =
car_inference::resource_policy::local_admission_for_scope(&root)
{
let owner = WorkerResidentOwner {
root,
allocation_id,
logical_model_id: model_id,
coordinator,
};
residents.push(owner.clone());
candidate_owner = Some(owner);
}
}
}
Self {
worker: Some(worker),
residents,
candidate: candidate_owner,
resident_models,
teardown_started: false,
}
}
fn worker_mut(&mut self) -> &mut WorkerProc {
self.worker.as_mut().expect("worker process owned")
}
fn charge_candidate(&self, measured_weights_bytes: u64) {
if let Some(candidate) = &self.candidate {
candidate
.coordinator
.mark_teardown_pending_allocation_with_charge(
&candidate.logical_model_id,
&candidate.allocation_id,
measured_weights_bytes,
);
}
}
fn clear_candidate(&mut self) {
let Some(candidate) = self.candidate.take() else {
return;
};
candidate
.coordinator
.finish_teardown_allocation(&candidate.logical_model_id, &candidate.allocation_id);
self.residents.retain(|resident| {
resident.root != candidate.root
|| resident.logical_model_id != candidate.logical_model_id
|| resident.allocation_id != candidate.allocation_id
});
}
fn charge_reported_model(
&mut self,
root: std::path::PathBuf,
model_id: &str,
allocation_id: String,
measured_weights_bytes: u64,
) {
let root = car_inference::resource_policy::normalized_state_root_key(&root);
if self
.residents
.iter()
.any(|resident| resident.root == root && resident.logical_model_id == model_id)
{
return;
}
let Some(coordinator) = car_inference::resource_policy::local_admission_for_scope(&root)
else {
return;
};
coordinator.mark_teardown_pending_allocation_with_charge(
model_id,
&allocation_id,
measured_weights_bytes,
);
self.residents.push(WorkerResidentOwner {
root,
logical_model_id: model_id.to_string(),
allocation_id,
coordinator,
});
}
fn begin_teardown(&mut self) {
if self.teardown_started {
return;
}
self.teardown_started = true;
for resident in &self.residents {
resident.coordinator.mark_teardown_pending_allocation(
&resident.logical_model_id,
&resident.allocation_id,
);
}
}
fn finish_accounting(
residents: &[WorkerResidentOwner],
resident_models: &std::sync::Mutex<WorkerResidentMap>,
) {
let mut tracked = resident_models
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
for resident in residents {
tracked.remove(&(resident.root.clone(), resident.logical_model_id.clone()));
resident
.coordinator
.finish_teardown_allocation(&resident.logical_model_id, &resident.allocation_id);
}
}
fn confirm_exited(mut self) {
self.worker.take();
Self::finish_accounting(&self.residents, &self.resident_models);
self.residents.clear();
}
fn return_to_slot(mut self, slot: &mut Option<WorkerProc>) {
*slot = self.worker.take();
self.residents.clear();
self.candidate = None;
}
async fn stop_and_confirm(mut self) -> Result<(), InferenceError> {
self.begin_teardown();
let worker = self.worker_mut();
match worker.child.try_wait().map_err(|error| {
InferenceError::InferenceFailed(format!(
"cannot inspect inference worker during teardown: {error}"
))
})? {
Some(_) => {}
None => {
worker.child.kill().await.map_err(|error| {
InferenceError::InferenceFailed(format!(
"cannot stop inference worker during teardown: {error}"
))
})?;
worker.child.wait().await.map_err(|error| {
InferenceError::InferenceFailed(format!(
"cannot reap inference worker during teardown: {error}"
))
})?;
}
}
self.confirm_exited();
Ok(())
}
}
impl Drop for WorkerProcessGuard {
fn drop(&mut self) {
self.begin_teardown();
let Some(mut worker) = self.worker.take() else {
return;
};
let residents = std::mem::take(&mut self.residents);
let resident_models = self.resident_models.clone();
let _ = worker.child.start_kill();
if tokio::runtime::Handle::try_current().is_ok() {
tokio::spawn(async move {
if worker.child.wait().await.is_ok() {
Self::finish_accounting(&residents, &resident_models);
}
});
} else {
std::mem::forget(worker);
}
}
}
const IDLE_REAP_INTERVAL: Duration = Duration::from_secs(60);
fn next_worker_allocation_scope() -> u64 {
static NEXT: AtomicU64 = AtomicU64::new(1);
NEXT.fetch_add(1, Ordering::Relaxed)
}
#[derive(Clone)]
pub struct WorkerOffload {
inner: Arc<Mutex<Option<WorkerProc>>>,
program: OsString,
args: Arc<Vec<OsString>>,
policy_generation: Arc<AtomicU64>,
resident_models: Arc<std::sync::Mutex<WorkerResidentMap>>,
allocation_scope: u64,
#[cfg(test)]
release_delay_ms: Arc<AtomicU64>,
}
impl WorkerOffload {
fn allocation_id(&self, model_id: &str) -> String {
format!("worker:{}:{model_id}", self.allocation_scope)
}
pub fn new() -> std::io::Result<Self> {
let exe = std::env::current_exe()?;
Ok(Self::with_command(
exe,
vec![OsString::from("--mlx-worker")],
))
}
pub fn with_command(program: impl Into<OsString>, args: Vec<OsString>) -> Self {
let me = Self {
inner: Arc::new(Mutex::new(None)),
program: program.into(),
args: Arc::new(args),
policy_generation: Arc::new(AtomicU64::new(1)),
resident_models: Arc::new(std::sync::Mutex::new(HashMap::new())),
allocation_scope: next_worker_allocation_scope(),
#[cfg(test)]
release_delay_ms: Arc::new(AtomicU64::new(0)),
};
me.spawn_idle_reaper();
me
}
fn spawn_idle_reaper(&self) {
if tokio::runtime::Handle::try_current().is_err() {
return;
}
let inner = Arc::clone(&self.inner);
let resident_models = Arc::clone(&self.resident_models);
tokio::spawn(async move {
let mut tick = tokio::time::interval(IDLE_REAP_INTERVAL);
tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
loop {
tick.tick().await;
if Arc::strong_count(&inner) == 1 {
let mut slot = inner.lock().await;
if let Some(worker) = slot.take() {
drop(WorkerProcessGuard::new(
worker,
resident_models.clone(),
None,
));
}
return;
}
let mut slot = inner.lock().await;
let Some(p) = slot.as_mut() else { continue };
match p.child.try_wait() {
Ok(Some(_)) => {
let dead = slot.take().expect("idle worker checked");
WorkerProcessGuard::new(dead, resident_models.clone(), None)
.confirm_exited();
tracing::info!("reaped a dead idle on-device inference worker");
}
Ok(None) => {}
Err(error) => {
tracing::warn!(%error, "cannot inspect idle inference worker; retaining ownership and residency");
}
}
}
});
}
fn spawn(&self) -> Result<WorkerProc, InferenceError> {
let mut cmd = Command::new(&self.program);
cmd.args(self.args.iter())
.env(WORKER_ENV, "1")
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::inherit())
.kill_on_drop(true);
let mut child = cmd.spawn().map_err(|e| {
InferenceError::InferenceFailed(format!("failed to spawn inference worker: {e}"))
})?;
let stdin = child
.stdin
.take()
.ok_or_else(|| InferenceError::InferenceFailed("worker stdin unavailable".into()))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| InferenceError::InferenceFailed("worker stdout unavailable".into()))?;
Ok(WorkerProc {
child,
stdin,
stdout: BufReader::new(stdout).lines(),
policy_generation: self.policy_generation.load(Ordering::Acquire),
state_root: None,
})
}
async fn take_or_spawn(
&self,
slot: &mut Option<WorkerProc>,
) -> Result<WorkerProc, InferenceError> {
if let Some(mut p) = slot.take() {
match p.child.try_wait() {
Ok(None)
if p.policy_generation == self.policy_generation.load(Ordering::Acquire) =>
{
return Ok(p)
}
Ok(Some(_)) => {
WorkerProcessGuard::new(p, self.resident_models.clone(), None).confirm_exited();
}
Ok(None) => {
WorkerProcessGuard::new(p, self.resident_models.clone(), None)
.stop_and_confirm()
.await?;
}
Err(error) => {
drop(WorkerProcessGuard::new(
p,
self.resident_models.clone(),
None,
));
return Err(InferenceError::InferenceFailed(format!(
"cannot inspect previous inference worker; teardown retained: {error}"
)));
}
}
}
self.spawn()
}
async fn take_or_spawn_for_scope(
&self,
slot: &mut Option<WorkerProc>,
state_root: &std::path::Path,
) -> Result<WorkerProc, InferenceError> {
let state_root = car_inference::resource_policy::normalized_state_root_key(state_root);
let mut worker = self.take_or_spawn(slot).await?;
if worker
.state_root
.as_ref()
.is_some_and(|current| current != &state_root)
{
WorkerProcessGuard::new(worker, self.resident_models.clone(), None)
.stop_and_confirm()
.await?;
worker = self.spawn()?;
}
worker.state_root = Some(state_root);
Ok(worker)
}
fn remember_residency(&self, residency: &WorkerResidencyAck, state_root: std::path::PathBuf) {
if residency.retention != car_inference::backend_cache::BackendRetention::Resident {
return;
}
let state_root = car_inference::resource_policy::normalized_state_root_key(&state_root);
let Some(coordinator) =
car_inference::resource_policy::local_admission_for_scope(&state_root)
else {
tracing::error!(root = %state_root.display(), model = %residency.model_id, "worker reported resident weights without a scoped admission owner");
return;
};
let allocation_id = self.allocation_id(&residency.model_id);
self.resident_models
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(
(state_root, residency.model_id.clone()),
WorkerResident {
allocation_id,
coordinator,
},
);
}
}
async fn write_line<W, T>(w: &mut W, value: &T) -> std::io::Result<()>
where
W: AsyncWrite + Unpin,
T: Serialize,
{
let mut line = serde_json::to_string(value).map_err(std::io::Error::other)?;
line.push('\n');
w.write_all(line.as_bytes()).await?;
w.flush().await
}
async fn do_generate(
proc: &mut WorkerProc,
request: GenerateRequest,
admission: LocalWorkerAdmission,
) -> Exchange<(InferenceResult, WorkerResidencyAck)> {
let req = WorkerRequest::Generate {
request: Box::new(request),
admission,
};
if let Err(e) = write_line(&mut proc.stdin, &req).await {
return Exchange::Dead(format!("write to inference worker failed: {e}"));
}
match proc.stdout.next_line().await {
Ok(Some(line)) => match serde_json::from_str::<WorkerResponse>(&line) {
Ok(WorkerResponse::Result { result, residency }) => Exchange::Ok((*result, residency)),
Ok(WorkerResponse::Error(msg)) => {
Exchange::Reported(InferenceError::InferenceFailed(msg))
}
Ok(WorkerResponse::LocalResourceBlocked {
preflight,
recovery,
}) => Exchange::Reported(InferenceError::LocalResourceBlocked {
preflight,
recovery,
}),
Ok(_) => Exchange::Dead("inference worker sent an unexpected response frame".into()),
Err(e) => Exchange::Dead(format!("inference worker sent invalid JSON: {e}")),
},
Ok(None) => Exchange::Dead("inference worker exited mid-request (EOF)".into()),
Err(e) => Exchange::Dead(format!("read from inference worker failed: {e}")),
}
}
#[async_trait::async_trait]
impl LocalGenerationOffload for WorkerOffload {
async fn generate(&self, _request: GenerateRequest) -> Result<InferenceResult, InferenceError> {
Err(InferenceError::InferenceFailed(
"WorkerOffload requires admission-aware dispatch".into(),
))
}
async fn stream(
&self,
_request: GenerateRequest,
) -> Result<tokio::sync::mpsc::Receiver<StreamEvent>, InferenceError> {
Err(InferenceError::InferenceFailed(
"WorkerOffload requires admission-aware dispatch".into(),
))
}
async fn generate_admitted(
&self,
request: GenerateRequest,
admission: LocalWorkerAdmission,
) -> Result<LocalOffloadResult, InferenceError> {
let mut slot = self.inner.lock().await;
let expected_model = request.model.clone().unwrap_or_default();
let candidate = (
admission.state_root.clone(),
expected_model.clone(),
self.allocation_id(&expected_model),
);
let proc = self
.take_or_spawn_for_scope(&mut slot, &admission.state_root)
.await?;
let mut ownership =
WorkerProcessGuard::new(proc, self.resident_models.clone(), Some(candidate));
ownership.charge_candidate(admission.measured_weights_bytes);
let state_root = admission.state_root.clone();
match do_generate(ownership.worker_mut(), request, admission).await {
Exchange::Ok((ir, residency)) => {
if residency.model_id != expected_model {
ownership.charge_reported_model(
state_root,
&residency.model_id,
self.allocation_id(&residency.model_id),
residency.measured_weights_bytes,
);
drop(ownership);
return Err(InferenceError::InferenceFailed(format!(
"local worker acknowledged model '{}' for requested '{}'",
residency.model_id, expected_model
)));
}
if residency.retention != car_inference::backend_cache::BackendRetention::Resident {
ownership.clear_candidate();
}
self.remember_residency(&residency, state_root);
ownership.return_to_slot(&mut slot);
Ok(LocalOffloadResult {
result: ir,
residency: LocalWorkerResidency {
model_id: residency.model_id,
measured_weights_bytes: residency.measured_weights_bytes,
},
retention: residency.retention,
})
}
Exchange::Reported(error) => {
if matches!(error, InferenceError::LocalResourceBlocked { .. }) {
ownership.clear_candidate();
ownership.return_to_slot(&mut slot);
} else {
drop(ownership);
}
Err(error)
}
Exchange::Dead(msg) => {
drop(ownership);
Err(InferenceError::InferenceFailed(format!(
"on-device inference worker crashed and was restarted; \
this request failed but the daemon is up: {msg}"
)))
}
}
}
async fn stream_admitted(
&self,
request: GenerateRequest,
admission: LocalWorkerAdmission,
) -> Result<LocalOffloadStream, InferenceError> {
let mut guard = self.inner.clone().lock_owned().await;
let expected_model = request.model.clone().unwrap_or_default();
let candidate = (
admission.state_root.clone(),
expected_model.clone(),
self.allocation_id(&expected_model),
);
let proc = self
.take_or_spawn_for_scope(&mut guard, &admission.state_root)
.await?;
let mut ownership =
WorkerProcessGuard::new(proc, self.resident_models.clone(), Some(candidate));
ownership.charge_candidate(admission.measured_weights_bytes);
let state_root = admission.state_root.clone();
{
let req = WorkerRequest::Stream {
request: Box::new(request),
admission,
};
if let Err(e) = write_line(&mut ownership.worker_mut().stdin, &req).await {
return Err(InferenceError::InferenceFailed(format!(
"write to inference worker failed: {e}"
)));
}
}
let residency = match ownership.worker_mut().stdout.next_line().await {
Ok(Some(line)) => match serde_json::from_str::<WorkerResponse>(&line) {
Ok(WorkerResponse::StreamStarted { residency }) => residency,
Ok(WorkerResponse::LocalResourceBlocked {
preflight,
recovery,
}) => {
ownership.clear_candidate();
ownership.return_to_slot(&mut guard);
return Err(InferenceError::LocalResourceBlocked {
preflight,
recovery,
});
}
Ok(WorkerResponse::Error(message)) => {
drop(ownership);
return Err(InferenceError::InferenceFailed(message));
}
Ok(_) => {
return Err(InferenceError::InferenceFailed(
"inference worker streamed before a successful load acknowledgement".into(),
));
}
Err(error) => {
return Err(InferenceError::InferenceFailed(format!(
"inference worker sent invalid stream acknowledgement: {error}"
)));
}
},
Ok(None) => {
return Err(InferenceError::InferenceFailed(
"inference worker exited before loading the streaming model".into(),
));
}
Err(error) => {
return Err(InferenceError::InferenceFailed(format!(
"failed reading inference worker load acknowledgement: {error}"
)));
}
};
if residency.model_id != expected_model {
ownership.charge_reported_model(
state_root,
&residency.model_id,
self.allocation_id(&residency.model_id),
residency.measured_weights_bytes,
);
drop(ownership);
return Err(InferenceError::InferenceFailed(format!(
"local worker acknowledged model '{}' for requested '{}'",
residency.model_id, expected_model
)));
}
if residency.retention != car_inference::backend_cache::BackendRetention::Resident {
ownership.clear_candidate();
}
self.remember_residency(&residency, state_root);
let (tx, rx) = tokio::sync::mpsc::channel::<StreamEvent>(64);
tokio::spawn(async move {
let mut clean = false;
loop {
match ownership.worker_mut().stdout.next_line().await {
Ok(Some(line)) => match serde_json::from_str::<WorkerResponse>(&line) {
Ok(WorkerResponse::Event(ev)) => {
if tx.send(*ev).await.is_err() {
break; }
}
Ok(WorkerResponse::StreamEnd) => {
clean = true;
break;
}
Ok(WorkerResponse::Error(msg)) => {
tracing::warn!(error = %msg, "inference worker stream error");
let _ = tx.send(StreamEvent::StopReason("error".into())).await;
break;
}
Ok(WorkerResponse::Result { .. })
| Ok(WorkerResponse::StreamStarted { .. })
| Ok(WorkerResponse::LocalResourceBlocked { .. })
| Err(_) => {
let _ = tx.send(StreamEvent::StopReason("error".into())).await;
break;
}
},
Ok(None) | Err(_) => {
let _ = tx.send(StreamEvent::StopReason("error".into())).await;
break;
}
}
}
if clean {
ownership.return_to_slot(&mut guard);
} else {
drop(ownership);
}
});
Ok(LocalOffloadStream {
events: rx,
residency: LocalWorkerResidency {
model_id: residency.model_id,
measured_weights_bytes: residency.measured_weights_bytes,
},
retention: residency.retention,
})
}
fn refresh_resource_policy(&self, generation: u64) {
self.policy_generation.store(generation, Ordering::Release);
if let Ok(mut slot) = self.inner.try_lock() {
if let Some(worker) = slot.take() {
drop(WorkerProcessGuard::new(
worker,
self.resident_models.clone(),
None,
));
}
}
}
fn resident_allocation_id(&self, model_id: &str) -> Option<String> {
Some(self.allocation_id(model_id))
}
async fn resident_models(&self) -> Vec<String> {
self.resident_models
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.keys()
.map(|(_, model_id)| model_id.clone())
.collect()
}
async fn release_model(&self, model_id: &str) -> Result<bool, InferenceError> {
let mut slot = self.inner.lock().await;
let resident = self
.resident_models
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.keys()
.any(|(_, resident_model)| resident_model == model_id);
if !resident {
return Ok(false);
}
let Some(worker) = slot.take() else {
return Err(InferenceError::InferenceFailed(format!(
"cannot confirm worker exit for resident model {model_id}: worker slot is empty"
)));
};
let mut release = WorkerProcessGuard::new(worker, self.resident_models.clone(), None);
release.begin_teardown();
#[cfg(test)]
tokio::time::sleep(Duration::from_millis(
self.release_delay_ms.load(Ordering::Acquire),
))
.await;
let worker = release.worker_mut();
match worker.child.try_wait().map_err(|error| {
InferenceError::InferenceFailed(format!(
"cannot inspect worker before releasing {model_id}: {error}"
))
})? {
Some(_) => {}
None => {
worker.child.kill().await.map_err(|error| {
InferenceError::InferenceFailed(format!(
"cannot stop worker before releasing {model_id}: {error}"
))
})?;
worker.child.wait().await.map_err(|error| {
InferenceError::InferenceFailed(format!(
"cannot reap worker before releasing {model_id}: {error}"
))
})?;
}
}
release.confirm_exited();
Ok(true)
}
}
pub async fn run_mlx_worker() {
std::env::set_var(WORKER_ENV, "1");
let mut engine: Option<(u64, std::path::PathBuf, Arc<InferenceEngine>)> = None;
let mut lines = BufReader::new(tokio::io::stdin()).lines();
let mut stdout = tokio::io::stdout();
while let Ok(Some(line)) = lines.next_line().await {
let line = line.trim();
if line.is_empty() {
continue;
}
let req: WorkerRequest = match serde_json::from_str(line) {
Ok(r) => r,
Err(e) => {
let _ = write_line(
&mut stdout,
&WorkerResponse::Error(format!("malformed worker request: {e}")),
)
.await;
continue;
}
};
let admission = match &req {
WorkerRequest::Generate { admission, .. } | WorkerRequest::Stream { admission, .. } => {
admission
}
};
let recreate = engine.as_ref().is_none_or(|(generation, root, _)| {
*generation != admission.policy_generation || root != &admission.state_root
});
if recreate {
let mut config = InferenceConfig::default();
config.state_root = admission.state_root.clone();
let candidate = Arc::new(InferenceEngine::new(config));
candidate.apply_local_resource_policy(admission.policy.clone());
engine = Some((
admission.policy_generation,
admission.state_root.clone(),
candidate,
));
}
let active_engine = Arc::clone(&engine.as_ref().expect("worker engine initialized").2);
active_engine.apply_local_resource_policy(admission.policy.clone());
match req {
WorkerRequest::Generate {
request,
admission: _,
} => {
let model_id = request.model.clone().unwrap_or_default();
let resp = match active_engine.generate_tracked(*request).await {
Ok(ir) => {
let measured = measured_worker_model_bytes(&active_engine, &model_id);
let retention = worker_model_retention(&active_engine, &model_id);
WorkerResponse::Result {
result: Box::new(ir),
residency: WorkerResidencyAck {
model_id,
measured_weights_bytes: measured,
retention,
},
}
}
Err(InferenceError::LocalResourceBlocked {
preflight,
recovery,
}) => WorkerResponse::LocalResourceBlocked {
preflight,
recovery,
},
Err(e) => WorkerResponse::Error(e.to_string()),
};
if write_line(&mut stdout, &resp).await.is_err() {
break; }
}
WorkerRequest::Stream {
request,
admission: _,
} => {
let model_id = request.model.clone().unwrap_or_default();
match active_engine.generate_tracked_stream(*request).await {
Ok(mut tracked) => {
let measured = measured_worker_model_bytes(&active_engine, &model_id);
let retention = worker_model_retention(&active_engine, &model_id);
if write_line(
&mut stdout,
&WorkerResponse::StreamStarted {
residency: WorkerResidencyAck {
model_id,
measured_weights_bytes: measured,
retention,
},
},
)
.await
.is_err()
{
break;
}
let mut broke = false;
while let Some(ev) = tracked.events.recv().await {
if write_line(&mut stdout, &WorkerResponse::Event(Box::new(ev)))
.await
.is_err()
{
broke = true;
break;
}
}
if broke {
break;
}
if write_line(&mut stdout, &WorkerResponse::StreamEnd)
.await
.is_err()
{
break;
}
}
Err(InferenceError::LocalResourceBlocked {
preflight,
recovery,
}) => {
if write_line(
&mut stdout,
&WorkerResponse::LocalResourceBlocked {
preflight,
recovery,
},
)
.await
.is_err()
{
break;
}
}
Err(e) => {
if write_line(&mut stdout, &WorkerResponse::Error(e.to_string()))
.await
.is_err()
{
break;
}
}
}
}
}
}
}
fn measured_worker_model_bytes(engine: &InferenceEngine, model_id: &str) -> u64 {
let Some(schema) = engine
.unified_registry
.get(model_id)
.or_else(|| engine.unified_registry.find_by_name(model_id))
else {
return 0;
};
car_inference::backend_cache::estimate_model_size(&engine.config.models_dir.join(&schema.name))
}
fn worker_model_retention(
engine: &InferenceEngine,
model_id: &str,
) -> car_inference::backend_cache::BackendRetention {
engine.local_model_retention(model_id)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn local_model_preflight_worker_rechecks_before_allocation() {
let source = include_str!("inference_worker.rs");
assert!(source.contains("LOCAL_ADMISSION_BOUNDARY:worker-side-allocation"));
assert!(source.contains("active_engine.generate_tracked(*request).await"));
}
const RESULT_LINE: &str = r#"{"Result":{"result":{"text":"pong","tool_calls":[],"trace_id":"t","model_used":"stub","latency_ms":1,"usage":null},"residency":{"model_id":"stub","measured_weights_bytes":1,"retention":"resident"}}}"#;
const TRANSIENT_RESULT_LINE: &str = r#"{"Result":{"result":{"text":"pong","tool_calls":[],"trace_id":"t","model_used":"stub","latency_ms":1,"usage":null},"residency":{"model_id":"stub","measured_weights_bytes":1,"retention":"transient"}}}"#;
const STREAM_STARTED_LINE: &str = r#"{"StreamStarted":{"residency":{"model_id":"stub","measured_weights_bytes":1,"retention":"resident"}}}"#;
fn test_admission() -> LocalWorkerAdmission {
LocalWorkerAdmission {
policy: car_inference::ResourcePolicy::custom_gb(8.0).unwrap(),
policy_generation: 1,
state_root: std::path::PathBuf::from("/tmp/car-worker-root"),
measured_weights_bytes: 1,
}
}
#[test]
fn envelope_serde_round_trips() {
let admission = car_inference::LocalWorkerAdmission {
policy: car_inference::ResourcePolicy::custom_gb(8.0).unwrap(),
policy_generation: 7,
state_root: std::path::PathBuf::from("/tmp/car-worker-root"),
measured_weights_bytes: 3 * 1024 * 1024 * 1024,
};
let req = WorkerRequest::Generate {
request: Box::new(GenerateRequest {
prompt: "hi".into(),
..Default::default()
}),
admission: admission.clone(),
};
let json = serde_json::to_string(&req).unwrap();
assert!(json.starts_with(r#"{"Generate":"#), "got {json}");
let decoded = serde_json::from_str::<WorkerRequest>(&json).unwrap();
let WorkerRequest::Generate {
admission: decoded_admission,
..
} = decoded
else {
panic!("wrong request variant")
};
assert_eq!(decoded_admission, admission);
assert!(matches!(
serde_json::from_str::<WorkerResponse>(RESULT_LINE).unwrap(),
WorkerResponse::Result { .. }
));
let ev = WorkerResponse::Event(Box::new(StreamEvent::TextDelta("hi".into())));
let ev_json = serde_json::to_string(&ev).unwrap();
assert_eq!(ev_json, r#"{"Event":{"TextDelta":"hi"}}"#);
assert_eq!(
serde_json::to_string(&WorkerResponse::StreamEnd).unwrap(),
r#""StreamEnd""#
);
let sr = StreamEvent::StopReason("length".into());
let sr2: StreamEvent = serde_json::from_str(&serde_json::to_string(&sr).unwrap()).unwrap();
assert!(matches!(sr2, StreamEvent::StopReason(s) if s == "length"));
}
#[test]
fn worker_protocol_preserves_structured_resource_rejection_and_residency_ack() {
let preflight = car_inference::LocalLoadPreflight {
model_id: "mlx/test".into(),
estimate: car_inference::ModelMemoryEstimate {
weights_mb: 6144,
runtime_overhead_mb: 256,
context_overhead_mb: 128,
transient_margin_mb: 256,
estimated_peak_mb: 6784,
evidence: car_inference::ModelResourceEvidence::FileSystemMeasured,
},
configured_ceiling_mb: 4096,
resident_model_mb: 0,
active_reservations_mb: 0,
estimated_incremental_mb: 6144,
accelerator_total_mb: None,
accelerator_resident_mb: None,
accelerator_incremental_mb: None,
live_available_mb: Some(8192),
emergency_reserve_mb: 4096,
verdict: car_inference::LocalLoadVerdict::ExceedsConfiguredCeiling,
};
let blocked = WorkerResponse::LocalResourceBlocked {
preflight: preflight.clone(),
recovery: "choose a smaller model".into(),
};
let decoded: WorkerResponse =
serde_json::from_str(&serde_json::to_string(&blocked).unwrap()).unwrap();
assert!(matches!(
decoded,
WorkerResponse::LocalResourceBlocked {
preflight: actual,
..
} if actual == preflight
));
let started = WorkerResponse::StreamStarted {
residency: WorkerResidencyAck {
model_id: "mlx/test".into(),
measured_weights_bytes: 3 * 1024 * 1024 * 1024,
retention: car_inference::backend_cache::BackendRetention::Resident,
},
};
assert!(matches!(
serde_json::from_str::<WorkerResponse>(&serde_json::to_string(&started).unwrap())
.unwrap(),
WorkerResponse::StreamStarted { .. }
));
}
#[cfg(unix)]
fn sh(script: &str) -> WorkerOffload {
WorkerOffload::with_command("sh", vec![OsString::from("-c"), OsString::from(script)])
}
#[cfg(unix)]
fn stub_request() -> GenerateRequest {
GenerateRequest {
model: Some("stub".into()),
..Default::default()
}
}
#[cfg(unix)]
#[tokio::test]
async fn generate_round_trips_and_reuses_the_worker() {
let off = sh(&format!(
"while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done"
));
let r1 = off
.generate_admitted(stub_request(), test_admission())
.await
.unwrap();
assert_eq!(r1.result.text, "pong");
let r2 = off
.generate_admitted(stub_request(), test_admission())
.await
.unwrap();
assert_eq!(r2.result.text, "pong");
}
#[cfg(unix)]
#[tokio::test]
async fn transient_worker_ack_never_creates_parent_residency() {
let off = sh(&format!(
"while IFS= read -r line; do printf '%s\\n' '{TRANSIENT_RESULT_LINE}'; done"
));
let result = off
.generate_admitted(stub_request(), test_admission())
.await
.unwrap();
assert_eq!(
result.retention,
car_inference::backend_cache::BackendRetention::Transient
);
assert!(off.resident_models().await.is_empty());
}
#[cfg(unix)]
#[tokio::test]
async fn generic_post_load_error_tears_down_before_releasing_pending_charge() {
let fixture = tempfile::tempdir().unwrap();
let admission = LocalWorkerAdmission {
state_root: fixture.path().join("state"),
measured_weights_bytes: 512 * 1024 * 1024,
..test_admission()
};
let coordinator = car_inference::resource_policy::scoped_local_admission(
&admission.state_root,
admission.policy.clone(),
car_inference::hardware::HardwareInfo::detect(),
);
let off = sh(
r#"while IFS= read -r line; do sleep 0.2; printf '%s\n' '{"Error":"failed after cache publication"}'; done"#,
);
let request = GenerateRequest {
model: Some("stub".into()),
..Default::default()
};
let caller = off.clone();
let exchange =
tokio::spawn(async move { caller.generate_admitted(request, admission).await });
tokio::time::sleep(Duration::from_millis(30)).await;
assert!(coordinator.teardown_pending("stub"));
assert_eq!(coordinator.resident_model_mb(), 512);
assert!(exchange.await.unwrap().is_err());
assert!(off.inner.lock().await.is_none());
tokio::time::timeout(Duration::from_secs(2), async {
while coordinator.teardown_pending("stub") {
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("worker must be reaped before its pending machine charge clears");
assert!(coordinator.resident_allocation_ids("stub").is_empty());
}
#[cfg(unix)]
#[tokio::test]
async fn mismatched_worker_ack_tears_down_requested_and_reported_allocations() {
let fixture = tempfile::tempdir().unwrap();
let admission = LocalWorkerAdmission {
state_root: fixture.path().join("state"),
measured_weights_bytes: 64 * 1024 * 1024,
..test_admission()
};
let coordinator = car_inference::resource_policy::scoped_local_admission(
&admission.state_root,
admission.policy.clone(),
car_inference::hardware::HardwareInfo::detect(),
);
let off = sh(&format!(
"while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done"
));
let request = GenerateRequest {
model: Some("requested-a".into()),
..Default::default()
};
let error = off
.generate_admitted(request, admission)
.await
.err()
.expect("mismatched ack must fail");
assert!(error.to_string().contains("acknowledged model 'stub'"));
assert!(off.inner.lock().await.is_none());
tokio::time::timeout(Duration::from_secs(2), async {
while coordinator.teardown_pending("requested-a")
|| coordinator.teardown_pending("stub")
{
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("mismatch teardown must reap both exact allocation candidates");
assert!(coordinator
.resident_allocation_ids("requested-a")
.is_empty());
assert!(coordinator.resident_allocation_ids("stub").is_empty());
}
#[cfg(unix)]
#[tokio::test]
async fn mismatched_stream_ack_never_returns_or_remembers_worker() {
let fixture = tempfile::tempdir().unwrap();
let admission = LocalWorkerAdmission {
state_root: fixture.path().join("state"),
measured_weights_bytes: 64 * 1024 * 1024,
..test_admission()
};
let coordinator = car_inference::resource_policy::scoped_local_admission(
&admission.state_root,
admission.policy.clone(),
car_inference::hardware::HardwareInfo::detect(),
);
let off = sh(&format!(
"while IFS= read -r line; do printf '%s\\n' '{STREAM_STARTED_LINE}'; done"
));
let request = GenerateRequest {
model: Some("requested-stream-a".into()),
..Default::default()
};
let error = off
.stream_admitted(request, admission)
.await
.err()
.expect("mismatched stream ack must fail");
assert!(error.to_string().contains("acknowledged model 'stub'"));
assert!(off.inner.lock().await.is_none());
tokio::time::timeout(Duration::from_secs(2), async {
while coordinator.teardown_pending("requested-stream-a")
|| coordinator.teardown_pending("stub")
{
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("stream mismatch teardown must reap both allocation candidates");
assert!(off.resident_models().await.is_empty());
}
#[cfg(unix)]
#[tokio::test]
async fn replacement_worker_generation_is_charged_separately_until_old_exit() {
let fixture = tempfile::tempdir().unwrap();
let admission = LocalWorkerAdmission {
state_root: fixture.path().join("state"),
measured_weights_bytes: 64 * 1024 * 1024,
..test_admission()
};
let coordinator = car_inference::resource_policy::scoped_local_admission(
&admission.state_root,
admission.policy.clone(),
car_inference::hardware::HardwareInfo::detect(),
);
let script = format!("while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done");
let first = sh(&script);
let second = sh(&script);
assert_ne!(first.allocation_id("stub"), second.allocation_id("stub"));
let mut first_reservation = coordinator
.reserve_measured_host("stub", 64 * 1024 * 1024, 0)
.unwrap();
let first_result = first
.generate_admitted(
GenerateRequest {
model: Some("stub".into()),
..Default::default()
},
admission.clone(),
)
.await
.unwrap();
first_reservation.publish_resident_weights_as(
&first.allocation_id("stub"),
first_result.residency.measured_weights_bytes,
);
drop(first_reservation);
let mut second_reservation = coordinator
.reserve_measured_host("stub", 64 * 1024 * 1024, 0)
.unwrap();
let second_result = second
.generate_admitted(
GenerateRequest {
model: Some("stub".into()),
..Default::default()
},
admission,
)
.await
.unwrap();
second_reservation.publish_resident_weights_as(
&second.allocation_id("stub"),
second_result.residency.measured_weights_bytes,
);
drop(second_reservation);
assert_eq!(coordinator.resident_allocation_ids("stub").len(), 2);
assert!(first.release_model("stub").await.unwrap());
assert_eq!(coordinator.resident_allocation_ids("stub").len(), 1);
assert!(second.release_model("stub").await.unwrap());
}
#[cfg(unix)]
#[tokio::test]
async fn policy_generation_change_restarts_existing_worker() {
let off = sh(r#"while IFS= read -r line; do
printf '{"Result":{"result":{"text":"pong","tool_calls":[],"trace_id":"t","model_used":"%s","latency_ms":1,"usage":null},"residency":{"model_id":"stub","measured_weights_bytes":1,"retention":"resident"}}}\n' "$$"
done"#);
let first = off
.generate_admitted(stub_request(), test_admission())
.await
.unwrap()
.result
.model_used;
off.refresh_resource_policy(2);
let mut updated = test_admission();
updated.policy_generation = 2;
let second = off
.generate_admitted(stub_request(), updated)
.await
.unwrap()
.result
.model_used;
assert_ne!(first, second, "policy refresh must replace the old child");
}
#[cfg(unix)]
#[tokio::test]
async fn targeted_release_waits_for_worker_exit_and_clears_residency() {
let fixture = tempfile::tempdir().unwrap();
let admission = LocalWorkerAdmission {
state_root: fixture.path().join("state"),
..test_admission()
};
let coordinator = car_inference::resource_policy::scoped_local_admission(
&admission.state_root,
admission.policy.clone(),
car_inference::hardware::HardwareInfo::detect(),
);
let off = sh(&format!(
"while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done"
));
off.generate_admitted(stub_request(), admission)
.await
.unwrap();
coordinator.mark_resident_allocation("stub", &off.allocation_id("stub"), 1);
assert_eq!(off.resident_models().await, vec!["stub".to_string()]);
assert!(off.release_model("stub").await.unwrap());
assert!(off.resident_models().await.is_empty());
assert!(!coordinator.is_resident("stub"));
assert!(off.inner.lock().await.is_none());
}
#[cfg(unix)]
#[tokio::test]
async fn cancelled_worker_release_keeps_accounting_until_background_reap() {
let fixture = tempfile::tempdir().unwrap();
let admission = LocalWorkerAdmission {
state_root: fixture.path().join("state"),
..test_admission()
};
let coordinator = car_inference::resource_policy::scoped_local_admission(
&admission.state_root,
admission.policy.clone(),
car_inference::hardware::HardwareInfo::detect(),
);
let off = sh(&format!(
"while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done"
));
off.generate_admitted(stub_request(), admission)
.await
.unwrap();
coordinator.mark_resident_allocation("stub", &off.allocation_id("stub"), 1);
off.release_delay_ms.store(250, Ordering::Release);
let release_owner = off.clone();
let release = tokio::spawn(async move { release_owner.release_model("stub").await });
tokio::time::sleep(Duration::from_millis(20)).await;
assert!(coordinator.teardown_pending("stub"));
release.abort();
let _ = release.await;
tokio::time::timeout(Duration::from_secs(2), async {
while !off.resident_models().await.is_empty() || coordinator.is_resident("stub") {
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("cancelled release guard must retain, kill, reap, and clear accounting");
assert!(off.inner.lock().await.is_none());
}
#[cfg(unix)]
#[tokio::test]
async fn worker_crash_fails_one_call_gracefully_then_respawns() {
let off = sh("exit 1");
let e1 = off
.generate_admitted(GenerateRequest::default(), test_admission())
.await
.err()
.expect("worker crash must fail");
assert!(
e1.to_string().contains("crashed and was restarted"),
"unexpected error: {e1}"
);
let e2 = off
.generate_admitted(GenerateRequest::default(), test_admission())
.await
.err()
.expect("worker crash must fail");
assert!(
e2.to_string().contains("crashed and was restarted"),
"unexpected error: {e2}"
);
}
#[cfg(unix)]
#[tokio::test]
async fn stream_forwards_events_then_closes() {
let off = sh(&format!(
"IFS= read -r line; \
printf '%s\\n' '{STREAM_STARTED_LINE}'; \
printf '%s\\n' '{{\"Event\":{{\"TextDelta\":\"he\"}}}}'; \
printf '%s\\n' '{{\"Event\":{{\"TextDelta\":\"llo\"}}}}'; \
printf '%s\\n' '\"StreamEnd\"'; \
while IFS= read -r l; do :; done"
));
let mut rx = off
.stream_admitted(stub_request(), test_admission())
.await
.unwrap()
.events;
let mut got = String::new();
while let Some(ev) = rx.recv().await {
if let StreamEvent::TextDelta(t) = ev {
got.push_str(&t);
}
}
assert_eq!(got, "hello");
}
#[cfg(unix)]
#[tokio::test]
async fn stream_worker_death_closes_with_error_end() {
let off = sh(&format!("IFS= read -r line; printf '%s\\n' '{STREAM_STARTED_LINE}'; printf '%s\\n' '{{\"Event\":{{\"TextDelta\":\"partial\"}}}}'"));
let mut rx = off
.stream_admitted(stub_request(), test_admission())
.await
.unwrap()
.events;
let mut saw_text = false;
let mut saw_error_end = false;
while let Some(ev) = rx.recv().await {
match ev {
StreamEvent::TextDelta(t) if t == "partial" => saw_text = true,
StreamEvent::StopReason(s) if s == "error" => saw_error_end = true,
_ => {}
}
}
assert!(saw_text, "should have forwarded the partial delta");
assert!(
saw_error_end,
"mid-stream death should surface an error end"
);
}
}