1use std::collections::HashMap;
34use std::ffi::OsString;
35use std::process::Stdio;
36use std::sync::atomic::{AtomicU64, Ordering};
37use std::sync::Arc;
38
39use std::time::Duration;
40
41use car_inference::stream::StreamEvent;
42use car_inference::tasks::generate::GenerateRequest;
43use car_inference::{
44 InferenceConfig, InferenceEngine, InferenceError, InferenceResult, InferenceTerminationAck,
45 LocalGenerationOffload, LocalLoadPreflight, LocalOffloadResult, LocalOffloadStream,
46 LocalWorkerAdmission, LocalWorkerResidency,
47};
48use serde::{Deserialize, Serialize};
49use tokio::io::{AsyncBufReadExt, AsyncWrite, AsyncWriteExt, BufReader, Lines};
50use tokio::process::{Child, ChildStdin, ChildStdout, Command};
51use tokio::sync::{mpsc, oneshot, Mutex};
52
53type WorkerCancelSender = mpsc::UnboundedSender<oneshot::Sender<bool>>;
54
55struct WorkerControlRegistration {
56 inference_id: String,
57 controls: Arc<std::sync::Mutex<HashMap<String, WorkerCancelSender>>>,
58}
59
60impl WorkerControlRegistration {
61 fn install(
62 controls: Arc<std::sync::Mutex<HashMap<String, WorkerCancelSender>>>,
63 inference_id: String,
64 ) -> (Self, mpsc::UnboundedReceiver<oneshot::Sender<bool>>) {
65 let (sender, receiver) = mpsc::unbounded_channel();
66 controls
67 .lock()
68 .unwrap_or_else(std::sync::PoisonError::into_inner)
69 .insert(inference_id.clone(), sender);
70 (
71 Self {
72 inference_id,
73 controls,
74 },
75 receiver,
76 )
77 }
78}
79
80impl Drop for WorkerControlRegistration {
81 fn drop(&mut self) {
82 self.controls
83 .lock()
84 .unwrap_or_else(std::sync::PoisonError::into_inner)
85 .remove(&self.inference_id);
86 }
87}
88
89pub const WORKER_ENV: &str = "CAR_INFERENCE_WORKER";
94
95#[derive(Serialize, Deserialize)]
97enum WorkerRequest {
98 Generate {
100 request: Box<GenerateRequest>,
101 admission: LocalWorkerAdmission,
102 },
103 Stream {
105 request: Box<GenerateRequest>,
106 admission: LocalWorkerAdmission,
107 },
108}
109
110#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
111struct WorkerResidencyAck {
112 model_id: String,
113 measured_weights_bytes: u64,
114 retention: car_inference::backend_cache::BackendRetention,
115}
116
117#[derive(Serialize, Deserialize)]
119enum WorkerResponse {
120 Result {
122 result: Box<InferenceResult>,
123 residency: WorkerResidencyAck,
124 },
125 StreamStarted {
126 residency: WorkerResidencyAck,
127 },
128 Event(Box<StreamEvent>),
130 StreamEnd,
132 Error(String),
136 LocalResourceBlocked {
137 preflight: LocalLoadPreflight,
138 recovery: String,
139 },
140}
141
142enum Exchange<T> {
145 Ok(T),
146 Reported(InferenceError),
148 Dead(String),
150}
151
152struct WorkerProc {
154 child: Child,
155 stdin: ChildStdin,
156 stdout: Lines<BufReader<ChildStdout>>,
157 policy_generation: u64,
158 state_root: Option<std::path::PathBuf>,
159}
160
161#[derive(Clone)]
162struct WorkerResident {
163 allocation_id: String,
164 coordinator: Arc<car_inference::resource_policy::LocalAdmissionCoordinator>,
165}
166
167type WorkerResidentMap = HashMap<(std::path::PathBuf, String), WorkerResident>;
168
169#[derive(Clone)]
170struct WorkerResidentOwner {
171 root: std::path::PathBuf,
172 logical_model_id: String,
173 allocation_id: String,
174 coordinator: Arc<car_inference::resource_policy::LocalAdmissionCoordinator>,
175}
176
177struct WorkerProcessGuard {
178 worker: Option<WorkerProc>,
179 residents: Vec<WorkerResidentOwner>,
180 candidate: Option<WorkerResidentOwner>,
181 resident_models: Arc<std::sync::Mutex<WorkerResidentMap>>,
182 teardown_started: bool,
183}
184
185impl WorkerProcessGuard {
186 fn new(
187 worker: WorkerProc,
188 resident_models: Arc<std::sync::Mutex<WorkerResidentMap>>,
189 candidate: Option<(std::path::PathBuf, String, String)>,
190 ) -> Self {
191 let mut residents = resident_models
192 .lock()
193 .unwrap_or_else(std::sync::PoisonError::into_inner)
194 .iter()
195 .map(|((root, logical_model_id), resident)| WorkerResidentOwner {
196 root: root.clone(),
197 logical_model_id: logical_model_id.clone(),
198 allocation_id: resident.allocation_id.clone(),
199 coordinator: resident.coordinator.clone(),
200 })
201 .collect::<Vec<_>>();
202 let mut candidate_owner = None;
203 if let Some((root, model_id, allocation_id)) =
204 candidate.filter(|(_, model_id, _)| !model_id.is_empty())
205 {
206 let root = car_inference::resource_policy::normalized_state_root_key(&root);
207 if !residents
208 .iter()
209 .any(|resident| resident.root == root && resident.logical_model_id == model_id)
210 {
211 if let Some(coordinator) =
212 car_inference::resource_policy::local_admission_for_scope(&root)
213 {
214 let owner = WorkerResidentOwner {
215 root,
216 allocation_id,
217 logical_model_id: model_id,
218 coordinator,
219 };
220 residents.push(owner.clone());
221 candidate_owner = Some(owner);
222 }
223 }
224 }
225 Self {
226 worker: Some(worker),
227 residents,
228 candidate: candidate_owner,
229 resident_models,
230 teardown_started: false,
231 }
232 }
233
234 fn worker_mut(&mut self) -> &mut WorkerProc {
235 self.worker.as_mut().expect("worker process owned")
236 }
237
238 fn charge_candidate(&self, measured_weights_bytes: u64) {
239 if let Some(candidate) = &self.candidate {
240 candidate
241 .coordinator
242 .mark_teardown_pending_allocation_with_charge(
243 &candidate.logical_model_id,
244 &candidate.allocation_id,
245 measured_weights_bytes,
246 );
247 }
248 }
249
250 fn clear_candidate(&mut self) {
251 let Some(candidate) = self.candidate.take() else {
252 return;
253 };
254 candidate
255 .coordinator
256 .finish_teardown_allocation(&candidate.logical_model_id, &candidate.allocation_id);
257 self.residents.retain(|resident| {
258 resident.root != candidate.root
259 || resident.logical_model_id != candidate.logical_model_id
260 || resident.allocation_id != candidate.allocation_id
261 });
262 }
263
264 fn charge_reported_model(
265 &mut self,
266 root: std::path::PathBuf,
267 model_id: &str,
268 allocation_id: String,
269 measured_weights_bytes: u64,
270 ) {
271 let root = car_inference::resource_policy::normalized_state_root_key(&root);
272 if self
273 .residents
274 .iter()
275 .any(|resident| resident.root == root && resident.logical_model_id == model_id)
276 {
277 return;
278 }
279 let Some(coordinator) = car_inference::resource_policy::local_admission_for_scope(&root)
280 else {
281 return;
282 };
283 coordinator.mark_teardown_pending_allocation_with_charge(
284 model_id,
285 &allocation_id,
286 measured_weights_bytes,
287 );
288 self.residents.push(WorkerResidentOwner {
289 root,
290 logical_model_id: model_id.to_string(),
291 allocation_id,
292 coordinator,
293 });
294 }
295
296 fn begin_teardown(&mut self) {
297 if self.teardown_started {
298 return;
299 }
300 self.teardown_started = true;
301 for resident in &self.residents {
302 resident.coordinator.mark_teardown_pending_allocation(
303 &resident.logical_model_id,
304 &resident.allocation_id,
305 );
306 }
307 }
308
309 fn finish_accounting(
310 residents: &[WorkerResidentOwner],
311 resident_models: &std::sync::Mutex<WorkerResidentMap>,
312 ) {
313 let mut tracked = resident_models
314 .lock()
315 .unwrap_or_else(std::sync::PoisonError::into_inner);
316 for resident in residents {
317 tracked.remove(&(resident.root.clone(), resident.logical_model_id.clone()));
318 resident
319 .coordinator
320 .finish_teardown_allocation(&resident.logical_model_id, &resident.allocation_id);
321 }
322 }
323
324 fn confirm_exited(mut self) {
325 self.worker.take();
326 Self::finish_accounting(&self.residents, &self.resident_models);
327 self.residents.clear();
328 }
329
330 fn return_to_slot(mut self, slot: &mut Option<WorkerProc>) {
331 *slot = self.worker.take();
332 self.residents.clear();
333 self.candidate = None;
334 }
335
336 async fn stop_and_confirm(mut self) -> Result<(), InferenceError> {
337 self.begin_teardown();
338 let worker = self.worker_mut();
339 match worker.child.try_wait().map_err(|error| {
340 InferenceError::InferenceFailed(format!(
341 "cannot inspect inference worker during teardown: {error}"
342 ))
343 })? {
344 Some(_) => {}
345 None => {
346 worker.child.kill().await.map_err(|error| {
347 InferenceError::InferenceFailed(format!(
348 "cannot stop inference worker during teardown: {error}"
349 ))
350 })?;
351 worker.child.wait().await.map_err(|error| {
352 InferenceError::InferenceFailed(format!(
353 "cannot reap inference worker during teardown: {error}"
354 ))
355 })?;
356 }
357 }
358 self.confirm_exited();
359 Ok(())
360 }
361}
362
363impl Drop for WorkerProcessGuard {
364 fn drop(&mut self) {
365 self.begin_teardown();
366 let Some(mut worker) = self.worker.take() else {
367 return;
368 };
369 let residents = std::mem::take(&mut self.residents);
370 let resident_models = self.resident_models.clone();
371 let _ = worker.child.start_kill();
372 if tokio::runtime::Handle::try_current().is_ok() {
373 tokio::spawn(async move {
374 if worker.child.wait().await.is_ok() {
375 Self::finish_accounting(&residents, &resident_models);
376 }
377 });
380 } else {
381 std::mem::forget(worker);
384 }
385 }
386}
387
388const IDLE_REAP_INTERVAL: Duration = Duration::from_secs(60);
393
394fn next_worker_allocation_scope() -> u64 {
395 static NEXT: AtomicU64 = AtomicU64::new(1);
396 NEXT.fetch_add(1, Ordering::Relaxed)
397}
398
399#[derive(Clone)]
402pub struct WorkerOffload {
403 inner: Arc<Mutex<Option<WorkerProc>>>,
404 program: OsString,
405 args: Arc<Vec<OsString>>,
406 policy_generation: Arc<AtomicU64>,
407 resident_models: Arc<std::sync::Mutex<WorkerResidentMap>>,
408 controls: Arc<std::sync::Mutex<HashMap<String, WorkerCancelSender>>>,
411 allocation_scope: u64,
412 #[cfg(test)]
413 release_delay_ms: Arc<AtomicU64>,
414}
415
416impl WorkerOffload {
417 fn allocation_id(&self, model_id: &str) -> String {
418 format!("worker:{}:{model_id}", self.allocation_scope)
419 }
420 pub fn new() -> std::io::Result<Self> {
423 let exe = std::env::current_exe()?;
424 Ok(Self::with_command(
425 exe,
426 vec![OsString::from("--mlx-worker")],
427 ))
428 }
429
430 pub fn with_command(program: impl Into<OsString>, args: Vec<OsString>) -> Self {
434 let me = Self {
435 inner: Arc::new(Mutex::new(None)),
436 program: program.into(),
437 args: Arc::new(args),
438 policy_generation: Arc::new(AtomicU64::new(1)),
439 resident_models: Arc::new(std::sync::Mutex::new(HashMap::new())),
440 controls: Arc::new(std::sync::Mutex::new(HashMap::new())),
441 allocation_scope: next_worker_allocation_scope(),
442 #[cfg(test)]
443 release_delay_ms: Arc::new(AtomicU64::new(0)),
444 };
445 me.spawn_idle_reaper();
446 me
447 }
448
449 fn spawn_idle_reaper(&self) {
460 if tokio::runtime::Handle::try_current().is_err() {
461 return;
464 }
465 let inner = Arc::clone(&self.inner);
466 let resident_models = Arc::clone(&self.resident_models);
467 tokio::spawn(async move {
468 let mut tick = tokio::time::interval(IDLE_REAP_INTERVAL);
469 tick.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay);
470 loop {
471 tick.tick().await;
472 if Arc::strong_count(&inner) == 1 {
476 let mut slot = inner.lock().await;
477 if let Some(worker) = slot.take() {
478 drop(WorkerProcessGuard::new(
479 worker,
480 resident_models.clone(),
481 None,
482 ));
483 }
484 return;
485 }
486 let mut slot = inner.lock().await;
487 let Some(p) = slot.as_mut() else { continue };
488 match p.child.try_wait() {
489 Ok(Some(_)) => {
490 let dead = slot.take().expect("idle worker checked");
491 WorkerProcessGuard::new(dead, resident_models.clone(), None)
492 .confirm_exited();
493 tracing::info!("reaped a dead idle on-device inference worker");
494 }
495 Ok(None) => {}
496 Err(error) => {
497 tracing::warn!(%error, "cannot inspect idle inference worker; retaining ownership and residency");
498 }
499 }
500 }
501 });
502 }
503
504 fn spawn(&self) -> Result<WorkerProc, InferenceError> {
505 let mut cmd = Command::new(&self.program);
506 cmd.args(self.args.iter())
507 .env(WORKER_ENV, "1")
508 .stdin(Stdio::piped())
509 .stdout(Stdio::piped())
510 .stderr(Stdio::inherit())
513 .kill_on_drop(true);
514 let mut child = cmd.spawn().map_err(|e| {
515 InferenceError::InferenceFailed(format!("failed to spawn inference worker: {e}"))
516 })?;
517 let stdin = child
518 .stdin
519 .take()
520 .ok_or_else(|| InferenceError::InferenceFailed("worker stdin unavailable".into()))?;
521 let stdout = child
522 .stdout
523 .take()
524 .ok_or_else(|| InferenceError::InferenceFailed("worker stdout unavailable".into()))?;
525 Ok(WorkerProc {
526 child,
527 stdin,
528 stdout: BufReader::new(stdout).lines(),
529 policy_generation: self.policy_generation.load(Ordering::Acquire),
530 state_root: None,
531 })
532 }
533
534 async fn take_or_spawn(
541 &self,
542 slot: &mut Option<WorkerProc>,
543 ) -> Result<WorkerProc, InferenceError> {
544 if let Some(mut p) = slot.take() {
545 match p.child.try_wait() {
551 Ok(None)
552 if p.policy_generation == self.policy_generation.load(Ordering::Acquire) =>
553 {
554 return Ok(p)
555 }
556 Ok(Some(_)) => {
557 WorkerProcessGuard::new(p, self.resident_models.clone(), None).confirm_exited();
558 }
559 Ok(None) => {
560 WorkerProcessGuard::new(p, self.resident_models.clone(), None)
561 .stop_and_confirm()
562 .await?;
563 }
564 Err(error) => {
565 drop(WorkerProcessGuard::new(
566 p,
567 self.resident_models.clone(),
568 None,
569 ));
570 return Err(InferenceError::InferenceFailed(format!(
571 "cannot inspect previous inference worker; teardown retained: {error}"
572 )));
573 }
574 }
575 }
576 self.spawn()
577 }
578
579 async fn take_or_spawn_for_scope(
580 &self,
581 slot: &mut Option<WorkerProc>,
582 state_root: &std::path::Path,
583 ) -> Result<WorkerProc, InferenceError> {
584 let state_root = car_inference::resource_policy::normalized_state_root_key(state_root);
585 let mut worker = self.take_or_spawn(slot).await?;
586 if worker
587 .state_root
588 .as_ref()
589 .is_some_and(|current| current != &state_root)
590 {
591 WorkerProcessGuard::new(worker, self.resident_models.clone(), None)
592 .stop_and_confirm()
593 .await?;
594 worker = self.spawn()?;
595 }
596 worker.state_root = Some(state_root);
597 Ok(worker)
598 }
599
600 fn remember_residency(&self, residency: &WorkerResidencyAck, state_root: std::path::PathBuf) {
601 if residency.retention != car_inference::backend_cache::BackendRetention::Resident {
602 return;
603 }
604 let state_root = car_inference::resource_policy::normalized_state_root_key(&state_root);
605 let Some(coordinator) =
606 car_inference::resource_policy::local_admission_for_scope(&state_root)
607 else {
608 tracing::error!(root = %state_root.display(), model = %residency.model_id, "worker reported resident weights without a scoped admission owner");
609 return;
610 };
611 let allocation_id = self.allocation_id(&residency.model_id);
612 self.resident_models
613 .lock()
614 .unwrap_or_else(std::sync::PoisonError::into_inner)
615 .insert(
616 (state_root, residency.model_id.clone()),
617 WorkerResident {
618 allocation_id,
619 coordinator,
623 },
624 );
625 }
626}
627
628async fn write_line<W, T>(w: &mut W, value: &T) -> std::io::Result<()>
629where
630 W: AsyncWrite + Unpin,
631 T: Serialize,
632{
633 let mut line = serde_json::to_string(value).map_err(std::io::Error::other)?;
634 line.push('\n');
635 w.write_all(line.as_bytes()).await?;
636 w.flush().await
637}
638
639async fn do_generate(
642 proc: &mut WorkerProc,
643 request: GenerateRequest,
644 admission: LocalWorkerAdmission,
645) -> Exchange<(InferenceResult, WorkerResidencyAck)> {
646 let req = WorkerRequest::Generate {
647 request: Box::new(request),
648 admission,
649 };
650 if let Err(e) = write_line(&mut proc.stdin, &req).await {
651 return Exchange::Dead(format!("write to inference worker failed: {e}"));
652 }
653 match proc.stdout.next_line().await {
654 Ok(Some(line)) => match serde_json::from_str::<WorkerResponse>(&line) {
655 Ok(WorkerResponse::Result { result, residency }) => Exchange::Ok((*result, residency)),
656 Ok(WorkerResponse::Error(msg)) => {
657 Exchange::Reported(InferenceError::InferenceFailed(msg))
658 }
659 Ok(WorkerResponse::LocalResourceBlocked {
660 preflight,
661 recovery,
662 }) => Exchange::Reported(InferenceError::LocalResourceBlocked {
663 preflight,
664 recovery,
665 }),
666 Ok(_) => Exchange::Dead("inference worker sent an unexpected response frame".into()),
667 Err(e) => Exchange::Dead(format!("inference worker sent invalid JSON: {e}")),
668 },
669 Ok(None) => Exchange::Dead("inference worker exited mid-request (EOF)".into()),
670 Err(e) => Exchange::Dead(format!("read from inference worker failed: {e}")),
671 }
672}
673
674#[async_trait::async_trait]
675impl LocalGenerationOffload for WorkerOffload {
676 async fn generate(&self, _request: GenerateRequest) -> Result<InferenceResult, InferenceError> {
677 Err(InferenceError::InferenceFailed(
678 "WorkerOffload requires admission-aware dispatch".into(),
679 ))
680 }
681
682 async fn stream(
683 &self,
684 _request: GenerateRequest,
685 ) -> Result<tokio::sync::mpsc::Receiver<StreamEvent>, InferenceError> {
686 Err(InferenceError::InferenceFailed(
687 "WorkerOffload requires admission-aware dispatch".into(),
688 ))
689 }
690
691 async fn generate_admitted(
692 &self,
693 request: GenerateRequest,
694 admission: LocalWorkerAdmission,
695 ) -> Result<LocalOffloadResult, InferenceError> {
696 let mut slot = self.inner.lock().await;
702 let expected_model = request.model.clone().unwrap_or_default();
703 let candidate = (
704 admission.state_root.clone(),
705 expected_model.clone(),
706 self.allocation_id(&expected_model),
707 );
708 let proc = self
709 .take_or_spawn_for_scope(&mut slot, &admission.state_root)
710 .await?;
711 let mut ownership =
712 WorkerProcessGuard::new(proc, self.resident_models.clone(), Some(candidate));
713 ownership.charge_candidate(admission.measured_weights_bytes);
714 let state_root = admission.state_root.clone();
715 enum GenerateSelection {
716 Exchange(Exchange<(InferenceResult, WorkerResidencyAck)>),
717 Cancel(oneshot::Sender<bool>),
718 }
719 let inference_id = car_inference::current_inference_control_id();
720 let termination = car_inference::current_controlled_termination_token();
721 let (control_registration, mut cancel_rx) = inference_id
722 .map(|id| WorkerControlRegistration::install(self.controls.clone(), id))
723 .unzip();
724 let selected = if let Some(cancel_rx) = cancel_rx.as_mut() {
725 let worker = ownership.worker_mut();
726 tokio::select! {
727 exchange = do_generate(worker, request, admission) => GenerateSelection::Exchange(exchange),
728 request = cancel_rx.recv() => match request {
729 Some(ack) => GenerateSelection::Cancel(ack),
730 None => unreachable!("control registration owns the sender"),
731 },
732 }
733 } else {
734 GenerateSelection::Exchange(
735 do_generate(ownership.worker_mut(), request, admission).await,
736 )
737 };
738 let exchange = match selected {
739 GenerateSelection::Exchange(exchange) => exchange,
740 GenerateSelection::Cancel(ack) => {
741 let confirmed = ownership.stop_and_confirm().await.is_ok();
742 if confirmed {
743 if let Some(token) = termination.as_ref() {
744 token.confirm_exact_backend_termination();
745 }
746 }
747 let _ = ack.send(confirmed);
748 drop(control_registration);
749 return Err(if confirmed {
750 InferenceError::ControlledTermination
751 } else {
752 InferenceError::InferenceFailed(
753 "isolated inference worker termination was not confirmed".into(),
754 )
755 });
756 }
757 };
758 drop(control_registration);
759 match exchange {
760 Exchange::Ok((ir, residency)) => {
761 if residency.model_id != expected_model {
762 ownership.charge_reported_model(
763 state_root,
764 &residency.model_id,
765 self.allocation_id(&residency.model_id),
766 residency.measured_weights_bytes,
767 );
768 drop(ownership);
769 return Err(InferenceError::InferenceFailed(format!(
770 "local worker acknowledged model '{}' for requested '{}'",
771 residency.model_id, expected_model
772 )));
773 }
774 if residency.retention != car_inference::backend_cache::BackendRetention::Resident {
775 ownership.clear_candidate();
776 }
777 self.remember_residency(&residency, state_root);
778 ownership.return_to_slot(&mut slot);
779 Ok(LocalOffloadResult {
780 result: ir,
781 residency: LocalWorkerResidency {
782 model_id: residency.model_id,
783 measured_weights_bytes: residency.measured_weights_bytes,
784 },
785 retention: residency.retention,
786 })
787 }
788 Exchange::Reported(error) => {
789 if matches!(error, InferenceError::LocalResourceBlocked { .. }) {
790 ownership.clear_candidate();
792 ownership.return_to_slot(&mut slot);
793 } else {
794 drop(ownership);
798 }
799 Err(error)
800 }
801 Exchange::Dead(msg) => {
802 drop(ownership);
806 Err(InferenceError::InferenceFailed(format!(
807 "on-device inference worker crashed and was restarted; \
808 this request failed but the daemon is up: {msg}"
809 )))
810 }
811 }
812 }
813
814 async fn stream_admitted(
815 &self,
816 request: GenerateRequest,
817 admission: LocalWorkerAdmission,
818 ) -> Result<LocalOffloadStream, InferenceError> {
819 let mut guard = self.inner.clone().lock_owned().await;
827 let expected_model = request.model.clone().unwrap_or_default();
828 let candidate = (
829 admission.state_root.clone(),
830 expected_model.clone(),
831 self.allocation_id(&expected_model),
832 );
833 let proc = self
834 .take_or_spawn_for_scope(&mut guard, &admission.state_root)
835 .await?;
836 let mut ownership =
837 WorkerProcessGuard::new(proc, self.resident_models.clone(), Some(candidate));
838 ownership.charge_candidate(admission.measured_weights_bytes);
839 let state_root = admission.state_root.clone();
840 let inference_id = car_inference::current_inference_control_id();
841 let termination = car_inference::current_controlled_termination_token();
842 let (control_registration, mut cancel_rx) = inference_id
843 .map(|id| WorkerControlRegistration::install(self.controls.clone(), id))
844 .unzip();
845 {
848 let req = WorkerRequest::Stream {
849 request: Box::new(request),
850 admission,
851 };
852 if let Err(e) = write_line(&mut ownership.worker_mut().stdin, &req).await {
853 return Err(InferenceError::InferenceFailed(format!(
854 "write to inference worker failed: {e}"
855 )));
856 }
857 }
858
859 enum StreamStartSelection {
860 Frame(std::io::Result<Option<String>>),
861 Cancel(oneshot::Sender<bool>),
862 }
863 let stream_start = if let Some(cancel_rx) = cancel_rx.as_mut() {
867 let worker = ownership.worker_mut();
868 tokio::select! {
869 frame = worker.stdout.next_line() => StreamStartSelection::Frame(frame),
870 request = cancel_rx.recv() => match request {
871 Some(ack) => StreamStartSelection::Cancel(ack),
872 None => unreachable!("control registration owns the sender"),
873 },
874 }
875 } else {
876 StreamStartSelection::Frame(ownership.worker_mut().stdout.next_line().await)
877 };
878 let stream_start = match stream_start {
879 StreamStartSelection::Frame(frame) => frame,
880 StreamStartSelection::Cancel(ack) => {
881 let confirmed = ownership.stop_and_confirm().await.is_ok();
882 if confirmed {
883 if let Some(token) = termination.as_ref() {
884 token.confirm_exact_backend_termination();
885 }
886 }
887 let _ = ack.send(confirmed);
888 return Err(if confirmed {
889 InferenceError::ControlledTermination
890 } else {
891 InferenceError::InferenceFailed(
892 "isolated inference worker termination before StreamStarted was not confirmed"
893 .into(),
894 )
895 });
896 }
897 };
898 let residency = match stream_start {
899 Ok(Some(line)) => match serde_json::from_str::<WorkerResponse>(&line) {
900 Ok(WorkerResponse::StreamStarted { residency }) => residency,
901 Ok(WorkerResponse::LocalResourceBlocked {
902 preflight,
903 recovery,
904 }) => {
905 ownership.clear_candidate();
906 ownership.return_to_slot(&mut guard);
907 return Err(InferenceError::LocalResourceBlocked {
908 preflight,
909 recovery,
910 });
911 }
912 Ok(WorkerResponse::Error(message)) => {
913 drop(ownership);
915 return Err(InferenceError::InferenceFailed(message));
916 }
917 Ok(_) => {
918 return Err(InferenceError::InferenceFailed(
919 "inference worker streamed before a successful load acknowledgement".into(),
920 ));
921 }
922 Err(error) => {
923 return Err(InferenceError::InferenceFailed(format!(
924 "inference worker sent invalid stream acknowledgement: {error}"
925 )));
926 }
927 },
928 Ok(None) => {
929 return Err(InferenceError::InferenceFailed(
930 "inference worker exited before loading the streaming model".into(),
931 ));
932 }
933 Err(error) => {
934 return Err(InferenceError::InferenceFailed(format!(
935 "failed reading inference worker load acknowledgement: {error}"
936 )));
937 }
938 };
939 if residency.model_id != expected_model {
940 ownership.charge_reported_model(
941 state_root,
942 &residency.model_id,
943 self.allocation_id(&residency.model_id),
944 residency.measured_weights_bytes,
945 );
946 drop(ownership);
947 return Err(InferenceError::InferenceFailed(format!(
948 "local worker acknowledged model '{}' for requested '{}'",
949 residency.model_id, expected_model
950 )));
951 }
952 if residency.retention != car_inference::backend_cache::BackendRetention::Resident {
953 ownership.clear_candidate();
954 }
955 self.remember_residency(&residency, state_root);
956
957 let (tx, rx) = tokio::sync::mpsc::channel::<StreamEvent>(64);
958 tokio::spawn(async move {
959 let _control_registration = control_registration;
960 let mut ownership = Some(ownership);
961 let mut clean = false;
962 loop {
963 enum StreamSelection {
964 Frame(std::io::Result<Option<String>>),
965 Cancel(oneshot::Sender<bool>),
966 }
967 let selected = if let Some(cancel_rx) = cancel_rx.as_mut() {
968 let worker = ownership
969 .as_mut()
970 .expect("worker ownership remains until terminal")
971 .worker_mut();
972 tokio::select! {
973 frame = worker.stdout.next_line() => StreamSelection::Frame(frame),
974 request = cancel_rx.recv() => match request {
975 Some(ack) => StreamSelection::Cancel(ack),
976 None => unreachable!("control registration owns the sender"),
977 },
978 }
979 } else {
980 StreamSelection::Frame(
981 ownership
982 .as_mut()
983 .expect("worker ownership remains until terminal")
984 .worker_mut()
985 .stdout
986 .next_line()
987 .await,
988 )
989 };
990 let frame = match selected {
991 StreamSelection::Frame(frame) => frame,
992 StreamSelection::Cancel(ack) => {
993 let confirmed = ownership
994 .take()
995 .expect("worker ownership remains until terminal")
996 .stop_and_confirm()
997 .await
998 .is_ok();
999 if confirmed {
1000 if let Some(token) = termination.as_ref() {
1001 token.confirm_exact_backend_termination();
1002 }
1003 }
1004 let _ = ack.send(confirmed);
1005 return;
1006 }
1007 };
1008 match frame {
1009 Ok(Some(line)) => match serde_json::from_str::<WorkerResponse>(&line) {
1010 Ok(WorkerResponse::Event(ev)) => {
1011 if tx.send(*ev).await.is_err() {
1012 break; }
1014 }
1015 Ok(WorkerResponse::StreamEnd) => {
1016 clean = true;
1017 break;
1018 }
1019 Ok(WorkerResponse::Error(msg)) => {
1020 tracing::warn!(error = %msg, "inference worker stream error");
1021 let _ = tx.send(StreamEvent::StopReason("error".into())).await;
1024 break;
1025 }
1026 Ok(WorkerResponse::Result { .. })
1027 | Ok(WorkerResponse::StreamStarted { .. })
1028 | Ok(WorkerResponse::LocalResourceBlocked { .. })
1029 | Err(_) => {
1030 let _ = tx.send(StreamEvent::StopReason("error".into())).await;
1032 break;
1033 }
1034 },
1035 Ok(None) | Err(_) => {
1036 let _ = tx.send(StreamEvent::StopReason("error".into())).await;
1038 break;
1039 }
1040 }
1041 }
1042 if clean {
1043 ownership
1045 .take()
1046 .expect("worker ownership remains on clean end")
1047 .return_to_slot(&mut guard);
1048 } else {
1049 drop(ownership.take());
1052 }
1053 });
1054 Ok(LocalOffloadStream {
1055 events: rx,
1056 residency: LocalWorkerResidency {
1057 model_id: residency.model_id,
1058 measured_weights_bytes: residency.measured_weights_bytes,
1059 },
1060 retention: residency.retention,
1061 })
1062 }
1063
1064 fn refresh_resource_policy(&self, generation: u64) {
1065 self.policy_generation.store(generation, Ordering::Release);
1066 if let Ok(mut slot) = self.inner.try_lock() {
1067 if let Some(worker) = slot.take() {
1068 drop(WorkerProcessGuard::new(
1069 worker,
1070 self.resident_models.clone(),
1071 None,
1072 ));
1073 }
1074 }
1075 }
1076
1077 fn resident_allocation_id(&self, model_id: &str) -> Option<String> {
1078 Some(self.allocation_id(model_id))
1079 }
1080
1081 async fn resident_models(&self) -> Vec<String> {
1082 self.resident_models
1083 .lock()
1084 .unwrap_or_else(std::sync::PoisonError::into_inner)
1085 .keys()
1086 .map(|(_, model_id)| model_id.clone())
1087 .collect()
1088 }
1089
1090 async fn release_model(&self, model_id: &str) -> Result<bool, InferenceError> {
1091 let mut slot = self.inner.lock().await;
1092 let resident = self
1093 .resident_models
1094 .lock()
1095 .unwrap_or_else(std::sync::PoisonError::into_inner)
1096 .keys()
1097 .any(|(_, resident_model)| resident_model == model_id);
1098 if !resident {
1099 return Ok(false);
1100 }
1101 let Some(worker) = slot.take() else {
1102 return Err(InferenceError::InferenceFailed(format!(
1103 "cannot confirm worker exit for resident model {model_id}: worker slot is empty"
1104 )));
1105 };
1106 let mut release = WorkerProcessGuard::new(worker, self.resident_models.clone(), None);
1107 release.begin_teardown();
1108 #[cfg(test)]
1109 tokio::time::sleep(Duration::from_millis(
1110 self.release_delay_ms.load(Ordering::Acquire),
1111 ))
1112 .await;
1113 let worker = release.worker_mut();
1114 match worker.child.try_wait().map_err(|error| {
1115 InferenceError::InferenceFailed(format!(
1116 "cannot inspect worker before releasing {model_id}: {error}"
1117 ))
1118 })? {
1119 Some(_) => {}
1120 None => {
1121 worker.child.kill().await.map_err(|error| {
1122 InferenceError::InferenceFailed(format!(
1123 "cannot stop worker before releasing {model_id}: {error}"
1124 ))
1125 })?;
1126 worker.child.wait().await.map_err(|error| {
1127 InferenceError::InferenceFailed(format!(
1128 "cannot reap worker before releasing {model_id}: {error}"
1129 ))
1130 })?;
1131 }
1132 }
1133 release.confirm_exited();
1134 Ok(true)
1135 }
1136
1137 async fn terminate_inference(&self, inference_id: &str) -> InferenceTerminationAck {
1138 let sender = self
1139 .controls
1140 .lock()
1141 .unwrap_or_else(std::sync::PoisonError::into_inner)
1142 .get(inference_id)
1143 .cloned();
1144 let Some(sender) = sender else {
1145 return InferenceTerminationAck::Unconfirmed;
1146 };
1147 let (ack_tx, ack_rx) = oneshot::channel();
1148 if sender.send(ack_tx).is_err() {
1149 return InferenceTerminationAck::Unconfirmed;
1150 }
1151 match ack_rx.await {
1152 Ok(true) => InferenceTerminationAck::Confirmed,
1153 Ok(false) | Err(_) => InferenceTerminationAck::Unconfirmed,
1154 }
1155 }
1156}
1157
1158pub async fn run_mlx_worker() {
1169 std::env::set_var(WORKER_ENV, "1");
1172
1173 let mut engine: Option<(u64, std::path::PathBuf, Arc<InferenceEngine>)> = None;
1174 let mut lines = BufReader::new(tokio::io::stdin()).lines();
1175 let mut stdout = tokio::io::stdout();
1176
1177 while let Ok(Some(line)) = lines.next_line().await {
1178 let line = line.trim();
1179 if line.is_empty() {
1180 continue;
1181 }
1182 let req: WorkerRequest = match serde_json::from_str(line) {
1183 Ok(r) => r,
1184 Err(e) => {
1185 let _ = write_line(
1186 &mut stdout,
1187 &WorkerResponse::Error(format!("malformed worker request: {e}")),
1188 )
1189 .await;
1190 continue;
1191 }
1192 };
1193 let admission = match &req {
1194 WorkerRequest::Generate { admission, .. } | WorkerRequest::Stream { admission, .. } => {
1195 admission
1196 }
1197 };
1198 let recreate = engine.as_ref().is_none_or(|(generation, root, _)| {
1199 *generation != admission.policy_generation || root != &admission.state_root
1200 });
1201 if recreate {
1202 let mut config = InferenceConfig::default();
1203 config.state_root = admission.state_root.clone();
1204 let candidate = Arc::new(InferenceEngine::new(config));
1205 candidate.apply_local_resource_policy(admission.policy.clone());
1206 engine = Some((
1207 admission.policy_generation,
1208 admission.state_root.clone(),
1209 candidate,
1210 ));
1211 }
1212 let active_engine = Arc::clone(&engine.as_ref().expect("worker engine initialized").2);
1213 active_engine.apply_local_resource_policy(admission.policy.clone());
1214
1215 match req {
1216 WorkerRequest::Generate {
1217 request,
1218 admission: _,
1219 } => {
1220 let model_id = request.model.clone().unwrap_or_default();
1226 let resp = match active_engine.generate_tracked(*request).await {
1227 Ok(ir) => {
1228 let measured = measured_worker_model_bytes(&active_engine, &model_id);
1229 let retention = worker_model_retention(&active_engine, &model_id);
1230 WorkerResponse::Result {
1231 result: Box::new(ir),
1232 residency: WorkerResidencyAck {
1233 model_id,
1234 measured_weights_bytes: measured,
1235 retention,
1236 },
1237 }
1238 }
1239 Err(InferenceError::LocalResourceBlocked {
1240 preflight,
1241 recovery,
1242 }) => WorkerResponse::LocalResourceBlocked {
1243 preflight,
1244 recovery,
1245 },
1246 Err(e) => WorkerResponse::Error(e.to_string()),
1247 };
1248 if write_line(&mut stdout, &resp).await.is_err() {
1249 break; }
1251 }
1252 WorkerRequest::Stream {
1253 request,
1254 admission: _,
1255 } => {
1256 let model_id = request.model.clone().unwrap_or_default();
1257 match active_engine.generate_tracked_stream(*request).await {
1258 Ok(mut tracked) => {
1259 let measured = measured_worker_model_bytes(&active_engine, &model_id);
1260 let retention = worker_model_retention(&active_engine, &model_id);
1261 if write_line(
1262 &mut stdout,
1263 &WorkerResponse::StreamStarted {
1264 residency: WorkerResidencyAck {
1265 model_id,
1266 measured_weights_bytes: measured,
1267 retention,
1268 },
1269 },
1270 )
1271 .await
1272 .is_err()
1273 {
1274 break;
1275 }
1276 let mut broke = false;
1277 while let Some(ev) = tracked.events.recv().await {
1278 if write_line(&mut stdout, &WorkerResponse::Event(Box::new(ev)))
1279 .await
1280 .is_err()
1281 {
1282 broke = true;
1283 break;
1284 }
1285 }
1286 if broke {
1287 break;
1288 }
1289 if write_line(&mut stdout, &WorkerResponse::StreamEnd)
1290 .await
1291 .is_err()
1292 {
1293 break;
1294 }
1295 }
1296 Err(InferenceError::LocalResourceBlocked {
1297 preflight,
1298 recovery,
1299 }) => {
1300 if write_line(
1301 &mut stdout,
1302 &WorkerResponse::LocalResourceBlocked {
1303 preflight,
1304 recovery,
1305 },
1306 )
1307 .await
1308 .is_err()
1309 {
1310 break;
1311 }
1312 }
1313 Err(e) => {
1314 if write_line(&mut stdout, &WorkerResponse::Error(e.to_string()))
1315 .await
1316 .is_err()
1317 {
1318 break;
1319 }
1320 }
1321 }
1322 }
1323 }
1324 }
1325}
1326
1327fn measured_worker_model_bytes(engine: &InferenceEngine, model_id: &str) -> u64 {
1328 let Some(schema) = engine
1329 .unified_registry
1330 .get(model_id)
1331 .or_else(|| engine.unified_registry.find_by_name(model_id))
1332 else {
1333 return 0;
1334 };
1335 car_inference::backend_cache::estimate_model_size(&engine.config.models_dir.join(&schema.name))
1336}
1337
1338fn worker_model_retention(
1339 engine: &InferenceEngine,
1340 model_id: &str,
1341) -> car_inference::backend_cache::BackendRetention {
1342 engine.local_model_retention(model_id)
1343}
1344
1345#[cfg(test)]
1346mod tests {
1347 use super::*;
1348
1349 #[test]
1350 fn local_model_preflight_worker_rechecks_before_allocation() {
1351 let source = include_str!("inference_worker.rs");
1352 assert!(source.contains("LOCAL_ADMISSION_BOUNDARY:worker-side-allocation"));
1353 assert!(source.contains("active_engine.generate_tracked(*request).await"));
1354 }
1355
1356 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"}}}"#;
1359 #[cfg(unix)]
1363 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"}}}"#;
1364 #[cfg(unix)]
1365 const STREAM_STARTED_LINE: &str = r#"{"StreamStarted":{"residency":{"model_id":"stub","measured_weights_bytes":1,"retention":"resident"}}}"#;
1366
1367 #[cfg(unix)]
1368 fn test_admission() -> LocalWorkerAdmission {
1369 LocalWorkerAdmission {
1370 policy: car_inference::ResourcePolicy::custom_gb(8.0).unwrap(),
1371 policy_generation: 1,
1372 state_root: std::path::PathBuf::from("/tmp/car-worker-root"),
1373 measured_weights_bytes: 1,
1374 }
1375 }
1376
1377 #[test]
1378 fn envelope_serde_round_trips() {
1379 let admission = car_inference::LocalWorkerAdmission {
1381 policy: car_inference::ResourcePolicy::custom_gb(8.0).unwrap(),
1382 policy_generation: 7,
1383 state_root: std::path::PathBuf::from("/tmp/car-worker-root"),
1384 measured_weights_bytes: 3 * 1024 * 1024 * 1024,
1385 };
1386 let req = WorkerRequest::Generate {
1387 request: Box::new(GenerateRequest {
1388 prompt: "hi".into(),
1389 ..Default::default()
1390 }),
1391 admission: admission.clone(),
1392 };
1393 let json = serde_json::to_string(&req).unwrap();
1394 assert!(json.starts_with(r#"{"Generate":"#), "got {json}");
1395 let decoded = serde_json::from_str::<WorkerRequest>(&json).unwrap();
1396 let WorkerRequest::Generate {
1397 admission: decoded_admission,
1398 ..
1399 } = decoded
1400 else {
1401 panic!("wrong request variant")
1402 };
1403 assert_eq!(decoded_admission, admission);
1404
1405 assert!(matches!(
1407 serde_json::from_str::<WorkerResponse>(RESULT_LINE).unwrap(),
1408 WorkerResponse::Result { .. }
1409 ));
1410 let ev = WorkerResponse::Event(Box::new(StreamEvent::TextDelta("hi".into())));
1411 let ev_json = serde_json::to_string(&ev).unwrap();
1412 assert_eq!(ev_json, r#"{"Event":{"TextDelta":"hi"}}"#);
1413 assert_eq!(
1414 serde_json::to_string(&WorkerResponse::StreamEnd).unwrap(),
1415 r#""StreamEnd""#
1416 );
1417 let sr = StreamEvent::StopReason("length".into());
1419 let sr2: StreamEvent = serde_json::from_str(&serde_json::to_string(&sr).unwrap()).unwrap();
1420 assert!(matches!(sr2, StreamEvent::StopReason(s) if s == "length"));
1421 }
1422
1423 #[test]
1424 fn worker_protocol_preserves_structured_resource_rejection_and_residency_ack() {
1425 let preflight = car_inference::LocalLoadPreflight {
1426 model_id: "mlx/test".into(),
1427 estimate: car_inference::ModelMemoryEstimate {
1428 weights_mb: 6144,
1429 runtime_overhead_mb: 256,
1430 context_overhead_mb: 128,
1431 transient_margin_mb: 256,
1432 estimated_peak_mb: 6784,
1433 evidence: car_inference::ModelResourceEvidence::FileSystemMeasured,
1434 },
1435 configured_ceiling_mb: 4096,
1436 resident_model_mb: 0,
1437 active_reservations_mb: 0,
1438 estimated_incremental_mb: 6144,
1439 accelerator_total_mb: None,
1440 accelerator_resident_mb: None,
1441 accelerator_incremental_mb: None,
1442 live_available_mb: Some(8192),
1443 emergency_reserve_mb: 4096,
1444 verdict: car_inference::LocalLoadVerdict::ExceedsConfiguredCeiling,
1445 };
1446 let blocked = WorkerResponse::LocalResourceBlocked {
1447 preflight: preflight.clone(),
1448 recovery: "choose a smaller model".into(),
1449 };
1450 let decoded: WorkerResponse =
1451 serde_json::from_str(&serde_json::to_string(&blocked).unwrap()).unwrap();
1452 assert!(matches!(
1453 decoded,
1454 WorkerResponse::LocalResourceBlocked {
1455 preflight: actual,
1456 ..
1457 } if actual == preflight
1458 ));
1459
1460 let started = WorkerResponse::StreamStarted {
1461 residency: WorkerResidencyAck {
1462 model_id: "mlx/test".into(),
1463 measured_weights_bytes: 3 * 1024 * 1024 * 1024,
1464 retention: car_inference::backend_cache::BackendRetention::Resident,
1465 },
1466 };
1467 assert!(matches!(
1468 serde_json::from_str::<WorkerResponse>(&serde_json::to_string(&started).unwrap())
1469 .unwrap(),
1470 WorkerResponse::StreamStarted { .. }
1471 ));
1472 }
1473
1474 #[cfg(unix)]
1478 fn sh(script: &str) -> WorkerOffload {
1479 WorkerOffload::with_command("sh", vec![OsString::from("-c"), OsString::from(script)])
1480 }
1481
1482 #[cfg(unix)]
1483 fn stub_request() -> GenerateRequest {
1484 GenerateRequest {
1485 model: Some("stub".into()),
1486 ..Default::default()
1487 }
1488 }
1489
1490 #[cfg(unix)]
1491 #[tokio::test]
1492 async fn generate_round_trips_and_reuses_the_worker() {
1493 let off = sh(&format!(
1496 "while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done"
1497 ));
1498 let r1 = off
1499 .generate_admitted(stub_request(), test_admission())
1500 .await
1501 .unwrap();
1502 assert_eq!(r1.result.text, "pong");
1503 let r2 = off
1504 .generate_admitted(stub_request(), test_admission())
1505 .await
1506 .unwrap();
1507 assert_eq!(r2.result.text, "pong");
1508 }
1509
1510 #[cfg(unix)]
1511 #[tokio::test]
1512 async fn exact_worker_cancel_confirms_only_after_kill_and_wait() {
1513 let off = sh("while IFS= read -r line; do sleep 60; done");
1514 let caller = off.clone();
1515 let generation = tokio::spawn(async move {
1516 car_inference::scope_inference_control_id("exact-worker-id".to_string(), async move {
1517 caller
1518 .generate_admitted(stub_request(), test_admission())
1519 .await
1520 })
1521 .await
1522 });
1523
1524 tokio::time::timeout(Duration::from_secs(2), async {
1525 loop {
1526 if off
1527 .controls
1528 .lock()
1529 .unwrap_or_else(std::sync::PoisonError::into_inner)
1530 .contains_key("exact-worker-id")
1531 {
1532 break;
1533 }
1534 tokio::task::yield_now().await;
1535 }
1536 })
1537 .await
1538 .expect("worker request must register its exact control ID");
1539
1540 assert_eq!(
1541 off.terminate_inference("exact-worker-id").await,
1542 InferenceTerminationAck::Confirmed
1543 );
1544 assert!(generation.await.unwrap().is_err());
1545 assert!(off.inner.lock().await.is_none());
1546 }
1547
1548 #[cfg(unix)]
1549 #[tokio::test]
1550 async fn transient_worker_ack_never_creates_parent_residency() {
1551 let off = sh(&format!(
1552 "while IFS= read -r line; do printf '%s\\n' '{TRANSIENT_RESULT_LINE}'; done"
1553 ));
1554 let result = off
1555 .generate_admitted(stub_request(), test_admission())
1556 .await
1557 .unwrap();
1558
1559 assert_eq!(
1560 result.retention,
1561 car_inference::backend_cache::BackendRetention::Transient
1562 );
1563 assert!(off.resident_models().await.is_empty());
1564 }
1565
1566 #[cfg(unix)]
1586 static LEDGER_GUARD: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
1587
1588 #[cfg(unix)]
1589 #[tokio::test]
1590 async fn generic_post_load_error_tears_down_before_releasing_pending_charge() {
1591 let _ledger = LEDGER_GUARD.lock().await;
1592 let fixture = tempfile::tempdir().unwrap();
1593 let admission = LocalWorkerAdmission {
1594 state_root: fixture.path().join("state"),
1595 measured_weights_bytes: 512 * 1024 * 1024,
1596 ..test_admission()
1597 };
1598 let coordinator = car_inference::resource_policy::scoped_local_admission(
1599 &admission.state_root,
1600 admission.policy.clone(),
1601 car_inference::hardware::HardwareInfo::detect(),
1602 );
1603 let resident_baseline_mb = coordinator.owned_resident_model_mb_for_testing();
1604 let off = sh(
1605 r#"while IFS= read -r line; do sleep 0.2; printf '%s\n' '{"Error":"failed after cache publication"}'; done"#,
1606 );
1607 let request = GenerateRequest {
1608 model: Some("stub".into()),
1609 ..Default::default()
1610 };
1611
1612 let baseline_mb = coordinator.resident_model_mb();
1619 let caller = off.clone();
1620 let exchange =
1621 tokio::spawn(async move { caller.generate_admitted(request, admission).await });
1622 tokio::time::sleep(Duration::from_millis(30)).await;
1623 assert!(coordinator.teardown_pending("stub"));
1624 assert_eq!(
1625 coordinator.resident_model_mb().saturating_sub(baseline_mb),
1626 512,
1627 "this worker's own 512 MB charge must still be held before release"
1628 );
1629 assert_eq!(
1630 coordinator.resident_allocation_ids("stub"),
1631 vec![off.allocation_id("stub")]
1632 );
1633 assert_eq!(
1634 coordinator
1635 .owned_resident_model_mb_for_testing()
1636 .checked_sub(resident_baseline_mb),
1637 Some(512)
1638 );
1639 assert!(exchange.await.unwrap().is_err());
1640 assert!(off.inner.lock().await.is_none());
1641
1642 tokio::time::timeout(Duration::from_secs(2), async {
1643 while coordinator.teardown_pending("stub") {
1644 tokio::time::sleep(Duration::from_millis(10)).await;
1645 }
1646 })
1647 .await
1648 .expect("worker must be reaped before its pending machine charge clears");
1649 assert!(coordinator.resident_allocation_ids("stub").is_empty());
1650 assert_eq!(
1651 coordinator.owned_resident_model_mb_for_testing(),
1652 resident_baseline_mb
1653 );
1654 }
1655
1656 #[cfg(unix)]
1657 #[tokio::test]
1658 async fn mismatched_worker_ack_tears_down_requested_and_reported_allocations() {
1659 let _ledger = LEDGER_GUARD.lock().await;
1660 let fixture = tempfile::tempdir().unwrap();
1661 let admission = LocalWorkerAdmission {
1662 state_root: fixture.path().join("state"),
1663 measured_weights_bytes: 64 * 1024 * 1024,
1664 ..test_admission()
1665 };
1666 let coordinator = car_inference::resource_policy::scoped_local_admission(
1667 &admission.state_root,
1668 admission.policy.clone(),
1669 car_inference::hardware::HardwareInfo::detect(),
1670 );
1671 let off = sh(&format!(
1672 "while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done"
1673 ));
1674 let request = GenerateRequest {
1675 model: Some("requested-a".into()),
1676 ..Default::default()
1677 };
1678
1679 let error = off
1680 .generate_admitted(request, admission)
1681 .await
1682 .err()
1683 .expect("mismatched ack must fail");
1684 assert!(error.to_string().contains("acknowledged model 'stub'"));
1685 assert!(off.inner.lock().await.is_none());
1686 tokio::time::timeout(Duration::from_secs(2), async {
1687 while coordinator.teardown_pending("requested-a")
1688 || coordinator.teardown_pending("stub")
1689 {
1690 tokio::time::sleep(Duration::from_millis(10)).await;
1691 }
1692 })
1693 .await
1694 .expect("mismatch teardown must reap both exact allocation candidates");
1695 assert!(coordinator
1696 .resident_allocation_ids("requested-a")
1697 .is_empty());
1698 assert!(coordinator.resident_allocation_ids("stub").is_empty());
1699 }
1700
1701 #[cfg(unix)]
1702 #[tokio::test]
1703 async fn mismatched_stream_ack_never_returns_or_remembers_worker() {
1704 let _ledger = LEDGER_GUARD.lock().await;
1705 let fixture = tempfile::tempdir().unwrap();
1706 let admission = LocalWorkerAdmission {
1707 state_root: fixture.path().join("state"),
1708 measured_weights_bytes: 64 * 1024 * 1024,
1709 ..test_admission()
1710 };
1711 let coordinator = car_inference::resource_policy::scoped_local_admission(
1712 &admission.state_root,
1713 admission.policy.clone(),
1714 car_inference::hardware::HardwareInfo::detect(),
1715 );
1716 let off = sh(&format!(
1717 "while IFS= read -r line; do printf '%s\\n' '{STREAM_STARTED_LINE}'; done"
1718 ));
1719 let request = GenerateRequest {
1720 model: Some("requested-stream-a".into()),
1721 ..Default::default()
1722 };
1723
1724 let error = off
1725 .stream_admitted(request, admission)
1726 .await
1727 .err()
1728 .expect("mismatched stream ack must fail");
1729 assert!(error.to_string().contains("acknowledged model 'stub'"));
1730 assert!(off.inner.lock().await.is_none());
1731 tokio::time::timeout(Duration::from_secs(2), async {
1732 while coordinator.teardown_pending("requested-stream-a")
1733 || coordinator.teardown_pending("stub")
1734 {
1735 tokio::time::sleep(Duration::from_millis(10)).await;
1736 }
1737 })
1738 .await
1739 .expect("stream mismatch teardown must reap both allocation candidates");
1740 assert!(off.resident_models().await.is_empty());
1741 }
1742
1743 #[cfg(unix)]
1744 #[tokio::test]
1745 async fn replacement_worker_generation_is_charged_separately_until_old_exit() {
1746 let _ledger = LEDGER_GUARD.lock().await;
1747 let fixture = tempfile::tempdir().unwrap();
1748 let admission = LocalWorkerAdmission {
1749 state_root: fixture.path().join("state"),
1750 measured_weights_bytes: 64 * 1024 * 1024,
1751 ..test_admission()
1752 };
1753 let coordinator = car_inference::resource_policy::scoped_local_admission(
1754 &admission.state_root,
1755 admission.policy.clone(),
1756 car_inference::hardware::HardwareInfo::detect(),
1757 );
1758 let script = format!("while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done");
1759 let first = sh(&script);
1760 let second = sh(&script);
1761 assert_ne!(first.allocation_id("stub"), second.allocation_id("stub"));
1762
1763 let mut first_reservation = coordinator
1764 .reserve_measured_host("stub", 64 * 1024 * 1024, 0)
1765 .unwrap();
1766 let first_result = first
1767 .generate_admitted(
1768 GenerateRequest {
1769 model: Some("stub".into()),
1770 ..Default::default()
1771 },
1772 admission.clone(),
1773 )
1774 .await
1775 .unwrap();
1776 first_reservation.publish_resident_weights_as(
1777 &first.allocation_id("stub"),
1778 first_result.residency.measured_weights_bytes,
1779 );
1780 drop(first_reservation);
1781
1782 let mut second_reservation = coordinator
1783 .reserve_measured_host("stub", 64 * 1024 * 1024, 0)
1784 .unwrap();
1785 let second_result = second
1786 .generate_admitted(
1787 GenerateRequest {
1788 model: Some("stub".into()),
1789 ..Default::default()
1790 },
1791 admission,
1792 )
1793 .await
1794 .unwrap();
1795 second_reservation.publish_resident_weights_as(
1796 &second.allocation_id("stub"),
1797 second_result.residency.measured_weights_bytes,
1798 );
1799 drop(second_reservation);
1800 assert_eq!(coordinator.resident_allocation_ids("stub").len(), 2);
1801
1802 assert!(first.release_model("stub").await.unwrap());
1803 assert_eq!(coordinator.resident_allocation_ids("stub").len(), 1);
1804 assert!(second.release_model("stub").await.unwrap());
1805 }
1806
1807 #[cfg(unix)]
1808 #[tokio::test]
1809 async fn policy_generation_change_restarts_existing_worker() {
1810 let off = sh(r#"while IFS= read -r line; do
1811printf '{"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' "$$"
1812 done"#);
1813 let first = off
1814 .generate_admitted(stub_request(), test_admission())
1815 .await
1816 .unwrap()
1817 .result
1818 .model_used;
1819 off.refresh_resource_policy(2);
1820 let mut updated = test_admission();
1821 updated.policy_generation = 2;
1822 let second = off
1823 .generate_admitted(stub_request(), updated)
1824 .await
1825 .unwrap()
1826 .result
1827 .model_used;
1828 assert_ne!(first, second, "policy refresh must replace the old child");
1829 }
1830
1831 #[cfg(unix)]
1832 #[tokio::test]
1833 async fn targeted_release_waits_for_worker_exit_and_clears_residency() {
1834 let _ledger = LEDGER_GUARD.lock().await;
1835 let fixture = tempfile::tempdir().unwrap();
1836 let admission = LocalWorkerAdmission {
1837 state_root: fixture.path().join("state"),
1838 ..test_admission()
1839 };
1840 let coordinator = car_inference::resource_policy::scoped_local_admission(
1841 &admission.state_root,
1842 admission.policy.clone(),
1843 car_inference::hardware::HardwareInfo::detect(),
1844 );
1845 let off = sh(&format!(
1846 "while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done"
1847 ));
1848 off.generate_admitted(stub_request(), admission)
1849 .await
1850 .unwrap();
1851 coordinator.mark_resident_allocation("stub", &off.allocation_id("stub"), 1);
1852 assert_eq!(off.resident_models().await, vec!["stub".to_string()]);
1853 assert!(off.release_model("stub").await.unwrap());
1854 assert!(off.resident_models().await.is_empty());
1855 assert!(!coordinator.is_resident("stub"));
1856 assert!(off.inner.lock().await.is_none());
1857 }
1858
1859 #[cfg(unix)]
1860 #[tokio::test]
1861 async fn cancelled_worker_release_keeps_accounting_until_background_reap() {
1862 let _ledger = LEDGER_GUARD.lock().await;
1863 let fixture = tempfile::tempdir().unwrap();
1864 let admission = LocalWorkerAdmission {
1865 state_root: fixture.path().join("state"),
1866 ..test_admission()
1867 };
1868 let coordinator = car_inference::resource_policy::scoped_local_admission(
1869 &admission.state_root,
1870 admission.policy.clone(),
1871 car_inference::hardware::HardwareInfo::detect(),
1872 );
1873 let off = sh(&format!(
1874 "while IFS= read -r line; do printf '%s\\n' '{RESULT_LINE}'; done"
1875 ));
1876 off.generate_admitted(stub_request(), admission)
1877 .await
1878 .unwrap();
1879 coordinator.mark_resident_allocation("stub", &off.allocation_id("stub"), 1);
1880 off.release_delay_ms.store(250, Ordering::Release);
1881
1882 let release_owner = off.clone();
1883 let release = tokio::spawn(async move { release_owner.release_model("stub").await });
1884 tokio::time::sleep(Duration::from_millis(20)).await;
1885 assert!(coordinator.teardown_pending("stub"));
1886 release.abort();
1887 let _ = release.await;
1888
1889 tokio::time::timeout(Duration::from_secs(2), async {
1890 while !off.resident_models().await.is_empty() || coordinator.is_resident("stub") {
1891 tokio::time::sleep(Duration::from_millis(10)).await;
1892 }
1893 })
1894 .await
1895 .expect("cancelled release guard must retain, kill, reap, and clear accounting");
1896 assert!(off.inner.lock().await.is_none());
1897 }
1898
1899 #[cfg(unix)]
1906 #[tokio::test]
1907 async fn a_dropped_worker_generation_is_killed_reaped_and_unaccounted() {
1908 let _ledger = LEDGER_GUARD.lock().await;
1909 let fixture = tempfile::tempdir().unwrap();
1910 let admission = LocalWorkerAdmission {
1911 state_root: fixture.path().join("state"),
1912 measured_weights_bytes: 256 * 1024 * 1024,
1913 ..test_admission()
1914 };
1915 let coordinator = car_inference::resource_policy::scoped_local_admission(
1916 &admission.state_root,
1917 admission.policy.clone(),
1918 car_inference::hardware::HardwareInfo::detect(),
1919 );
1920 let resident_baseline_mb = coordinator.owned_resident_model_mb_for_testing();
1921 let pid_file = fixture.path().join("worker.pid");
1925 let off = sh(&format!(
1926 "IFS= read -r line; echo $$ > '{tmp}'; mv '{tmp}' '{pid}'; exec sleep 60",
1927 tmp = pid_file.with_extension("tmp").display(),
1928 pid = pid_file.display(),
1929 ));
1930 let alive = |pid: u32| {
1931 std::process::Command::new("sh")
1932 .arg("-c")
1933 .arg(format!("kill -0 {pid} 2>/dev/null"))
1934 .status()
1935 .is_ok_and(|status| status.success())
1936 };
1937
1938 let mut generation = Box::pin(off.generate_admitted(stub_request(), admission));
1939 let pid = tokio::time::timeout(Duration::from_secs(5), async {
1940 tokio::select! {
1941 result = &mut generation => {
1942 panic!("the stub never answers, yet generation ended (ok: {})", result.is_ok())
1943 }
1944 pid = async {
1945 loop {
1946 if let Some(pid) = std::fs::read_to_string(&pid_file)
1947 .ok()
1948 .and_then(|text| text.trim().parse::<u32>().ok())
1949 {
1950 break pid;
1951 }
1952 tokio::time::sleep(Duration::from_millis(10)).await;
1953 }
1954 } => pid,
1955 }
1956 })
1957 .await
1958 .expect("the stub worker must receive the request");
1959 assert!(
1960 alive(pid),
1961 "positive control: the in-flight worker is running"
1962 );
1963 assert!(
1964 coordinator.teardown_pending("stub"),
1965 "positive control: the in-flight request holds its admission charge"
1966 );
1967
1968 assert!(
1971 tokio::time::timeout(Duration::from_millis(50), generation)
1972 .await
1973 .is_err(),
1974 "a stub that never answers must hit the deadline"
1975 );
1976
1977 tokio::time::timeout(Duration::from_secs(5), async {
1978 while alive(pid) || coordinator.teardown_pending("stub") {
1979 tokio::time::sleep(Duration::from_millis(10)).await;
1980 }
1981 })
1982 .await
1983 .expect("a dropped generation must kill and reap its worker and clear its charge");
1984 assert!(coordinator.resident_allocation_ids("stub").is_empty());
1985 assert_eq!(
1986 coordinator.owned_resident_model_mb_for_testing(),
1987 resident_baseline_mb
1988 );
1989 assert!(off.resident_models().await.is_empty());
1990 assert!(
1991 off.inner.lock().await.is_none(),
1992 "a killed worker must not return to the slot"
1993 );
1994 }
1995
1996 #[cfg(unix)]
1997 #[tokio::test]
1998 async fn worker_crash_fails_one_call_gracefully_then_respawns() {
1999 let off = sh("exit 1");
2006 let e1 = off
2007 .generate_admitted(GenerateRequest::default(), test_admission())
2008 .await
2009 .err()
2010 .expect("worker crash must fail");
2011 assert!(
2012 e1.to_string().contains("crashed and was restarted"),
2013 "unexpected error: {e1}"
2014 );
2015 let e2 = off
2016 .generate_admitted(GenerateRequest::default(), test_admission())
2017 .await
2018 .err()
2019 .expect("worker crash must fail");
2020 assert!(
2021 e2.to_string().contains("crashed and was restarted"),
2022 "unexpected error: {e2}"
2023 );
2024 }
2025
2026 #[cfg(unix)]
2027 #[tokio::test]
2028 async fn stream_forwards_events_then_closes() {
2029 let off = sh(&format!(
2031 "IFS= read -r line; \
2032 printf '%s\\n' '{STREAM_STARTED_LINE}'; \
2033 printf '%s\\n' '{{\"Event\":{{\"TextDelta\":\"he\"}}}}'; \
2034 printf '%s\\n' '{{\"Event\":{{\"TextDelta\":\"llo\"}}}}'; \
2035 printf '%s\\n' '\"StreamEnd\"'; \
2036 while IFS= read -r l; do :; done"
2037 ));
2038 let mut rx = off
2039 .stream_admitted(stub_request(), test_admission())
2040 .await
2041 .unwrap()
2042 .events;
2043 let mut got = String::new();
2044 while let Some(ev) = rx.recv().await {
2045 if let StreamEvent::TextDelta(t) = ev {
2046 got.push_str(&t);
2047 }
2048 }
2049 assert_eq!(got, "hello");
2050 }
2051
2052 #[cfg(unix)]
2053 #[tokio::test]
2054 async fn stream_worker_death_closes_with_error_end() {
2055 let off = sh(&format!("IFS= read -r line; printf '%s\\n' '{STREAM_STARTED_LINE}'; printf '%s\\n' '{{\"Event\":{{\"TextDelta\":\"partial\"}}}}'"));
2058 let mut rx = off
2059 .stream_admitted(stub_request(), test_admission())
2060 .await
2061 .unwrap()
2062 .events;
2063 let mut saw_text = false;
2064 let mut saw_error_end = false;
2065 while let Some(ev) = rx.recv().await {
2066 match ev {
2067 StreamEvent::TextDelta(t) if t == "partial" => saw_text = true,
2068 StreamEvent::StopReason(s) if s == "error" => saw_error_end = true,
2069 _ => {}
2070 }
2071 }
2072 assert!(saw_text, "should have forwarded the partial delta");
2073 assert!(
2074 saw_error_end,
2075 "mid-stream death should surface an error end"
2076 );
2077 }
2078}