1use std::path::Path;
20use std::sync::Arc;
21use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
22use std::thread::JoinHandle;
23use std::time::Duration;
24
25use crate::failover::FailoverWatchdog;
26use crate::heartbeat::{HeartbeatError, HeartbeatTable};
27use crate::message_transport::{MessageTransport, TransportError};
28use crate::pass_registry::{execute as exec_pass, Pass, PassResult};
29use crate::shared_ring::{RingError, SharedRing, PAYLOAD_BYTES};
30
31const RESULT_TOKEN_OFFSET: usize = 4;
43const ARG_LEN_OFFSET: usize = 8;
44const ARGS_OFFSET: usize = 10;
45const MAX_ARG_LEN: usize = PAYLOAD_BYTES - ARGS_OFFSET;
46
47const RESULT_TOKEN_R_OFFSET: usize = 0;
52const STATUS_OFFSET: usize = 4;
53const RESULT_LEN_OFFSET: usize = 5;
54const RESULT_DATA_OFFSET: usize = 7;
55const MAX_RESULT_LEN: usize = PAYLOAD_BYTES - RESULT_DATA_OFFSET;
56
57#[derive(Debug, Clone, Copy, PartialEq, Eq)]
58pub enum SchedError {
59 Ring(RingError),
60 Transport(TransportError),
61 Heartbeat(HeartbeatError),
62 ArgsTooLarge,
63 ResultTooLarge,
64 UnsupportedTransportFamily(crate::mmf_dispatcher::MmfFamily),
68}
69
70impl From<RingError> for SchedError {
71 fn from(e: RingError) -> Self { Self::Ring(e) }
72}
73impl From<TransportError> for SchedError {
74 fn from(e: TransportError) -> Self { Self::Transport(e) }
75}
76impl From<HeartbeatError> for SchedError {
77 fn from(e: HeartbeatError) -> Self { Self::Heartbeat(e) }
78}
79
80pub struct Submitter {
86 submit_ring: Arc<dyn MessageTransport>,
87 next_token: AtomicU64,
88}
89
90impl Submitter {
91 pub fn new(submit_ring: Arc<dyn MessageTransport>) -> Self {
92 Self { submit_ring, next_token: AtomicU64::new(1) }
93 }
94
95 pub fn submit(&self, pass: &Pass) -> Result<u32, SchedError> {
97 if pass.args.len() > MAX_ARG_LEN {
98 return Err(SchedError::ArgsTooLarge);
99 }
100 let token = self.next_token.fetch_add(1, Ordering::Relaxed) as u32;
101 let mut slot = [0u8; PAYLOAD_BYTES];
102 slot[0..4].copy_from_slice(&pass.closure_id.to_le_bytes());
103 slot[RESULT_TOKEN_OFFSET..ARG_LEN_OFFSET]
104 .copy_from_slice(&token.to_le_bytes());
105 slot[ARG_LEN_OFFSET..ARGS_OFFSET]
106 .copy_from_slice(&(pass.args.len() as u16).to_le_bytes());
107 slot[ARGS_OFFSET..ARGS_OFFSET + pass.args.len()]
108 .copy_from_slice(&pass.args);
109 self.submit_ring.try_push(&slot)?;
110 Ok(token)
111 }
112}
113
114pub struct ResultCollector {
116 result_ring: Arc<dyn MessageTransport>,
117}
118
119#[derive(Debug, Clone)]
120pub struct SubmittedResult {
121 pub token: u32,
122 pub result: PassResult,
123}
124
125impl ResultCollector {
126 pub fn new(result_ring: Arc<dyn MessageTransport>) -> Self {
127 Self { result_ring }
128 }
129
130 pub fn try_recv(&self) -> Result<SubmittedResult, SchedError> {
133 let mut slot = [0u8; PAYLOAD_BYTES];
134 self.result_ring.try_pop(&mut slot)?;
135 Ok(parse_result_slot(&slot))
136 }
137}
138
139fn parse_result_slot(slot: &[u8; PAYLOAD_BYTES]) -> SubmittedResult {
140 let token = u32::from_le_bytes(
141 slot[RESULT_TOKEN_R_OFFSET..STATUS_OFFSET].try_into().unwrap()
142 );
143 let status = slot[STATUS_OFFSET];
144 let len = u16::from_le_bytes(
145 slot[RESULT_LEN_OFFSET..RESULT_DATA_OFFSET].try_into().unwrap()
146 ) as usize;
147 let data = slot[RESULT_DATA_OFFSET..RESULT_DATA_OFFSET + len.min(MAX_RESULT_LEN)].to_vec();
148 let result = if status == 0 {
149 Ok(data)
150 } else {
151 let msg = String::from_utf8_lossy(&data).into_owned();
152 Err(crate::pass_registry::PassError::ExecutionError(msg))
153 };
154 SubmittedResult { token, result }
155}
156
157pub struct BackgroundScheduler {
165 submit_ring: Arc<dyn MessageTransport>,
166 result_ring: Arc<dyn MessageTransport>,
167 heartbeat: Arc<HeartbeatTable>,
168 slot_idx: usize,
169 stop: Arc<AtomicBool>,
170 worker: Option<JoinHandle<()>>,
171 header_sidecar: subetha_core::HandshakeHeader,
172 ring_sidecar: Box<subetha_core::ObservationRing>,
173}
174
175impl subetha_sidecar::AdaptiveInstance for BackgroundScheduler {
176 fn header(&self) -> &subetha_core::HandshakeHeader { &self.header_sidecar }
177 fn ring(&self) -> &subetha_core::ObservationRing { &self.ring_sidecar }
178 fn make_policy(&self) -> Box<dyn subetha_sidecar::Policy> {
179 Box::new(subetha_sidecar::NoMigrationPolicy)
180 }
181}
182
183impl BackgroundScheduler {
184 pub fn start(
188 submit_path: impl AsRef<Path>,
189 result_path: impl AsRef<Path>,
190 heartbeat_path: impl AsRef<Path>,
191 capacity: usize,
192 heartbeat_capacity: usize,
193 ) -> Result<Self, SchedError> {
194 let submit_ring: Arc<dyn MessageTransport> = Arc::new(
195 SharedRing::open(submit_path.as_ref(), capacity)
196 .or_else(|_| SharedRing::create(submit_path.as_ref(), capacity))?
197 );
198 let result_ring: Arc<dyn MessageTransport> = Arc::new(
199 SharedRing::open(result_path.as_ref(), capacity)
200 .or_else(|_| SharedRing::create(result_path.as_ref(), capacity))?
201 );
202 Self::start_with_transports(
203 submit_ring,
204 result_ring,
205 heartbeat_path,
206 heartbeat_capacity,
207 )
208 }
209
210 pub fn start_with_transports(
216 submit_ring: Arc<dyn MessageTransport>,
217 result_ring: Arc<dyn MessageTransport>,
218 heartbeat_path: impl AsRef<Path>,
219 heartbeat_capacity: usize,
220 ) -> Result<Self, SchedError> {
221 let heartbeat = Arc::new(
222 HeartbeatTable::open(heartbeat_path.as_ref(), heartbeat_capacity)
223 .or_else(|_| HeartbeatTable::create(heartbeat_path.as_ref(), heartbeat_capacity))?
224 );
225 let pid = std::process::id();
226 let slot_idx = heartbeat.register(pid)?;
227
228 let stop = Arc::new(AtomicBool::new(false));
229 let stop_w = stop.clone();
230 let submit_w = submit_ring.clone();
231 let result_w = result_ring.clone();
232 let hb_w = heartbeat.clone();
233 let worker = std::thread::Builder::new()
234 .name(format!("subetha-sched-worker-pid{pid}"))
235 .spawn(move || {
236 Self::worker_loop(submit_w, result_w, hb_w, slot_idx, stop_w);
237 })
238 .expect("spawn scheduler worker thread");
239
240 Ok(Self {
241 submit_ring,
242 result_ring,
243 heartbeat,
244 slot_idx,
245 stop,
246 header_sidecar: subetha_core::HandshakeHeader::new(),
247 ring_sidecar: Box::new(subetha_core::ObservationRing::new()),
248 worker: Some(worker),
249 })
250 }
251
252 pub fn start_by_workload_shape(
267 submit_path: impl AsRef<Path>,
268 submit_shape: crate::mmf_dispatcher::MmfWorkloadShape,
269 result_path: impl AsRef<Path>,
270 result_shape: crate::mmf_dispatcher::MmfWorkloadShape,
271 heartbeat_path: impl AsRef<Path>,
272 capacity: usize,
273 heartbeat_capacity: usize,
274 ) -> Result<
275 (
276 Self,
277 crate::mmf_dispatcher::MmfFamily,
278 crate::mmf_dispatcher::MmfFamily,
279 ),
280 SchedError,
281 > {
282 let submit_family = crate::mmf_dispatcher::MmfDispatcher::pick(submit_shape);
283 let result_family = crate::mmf_dispatcher::MmfDispatcher::pick(result_shape);
284 let submit_transport =
285 Self::build_transport_for(submit_family, submit_path.as_ref(), capacity)?;
286 let result_transport =
287 Self::build_transport_for(result_family, result_path.as_ref(), capacity)?;
288 let sched = Self::start_with_transports(
289 submit_transport,
290 result_transport,
291 heartbeat_path,
292 heartbeat_capacity,
293 )?;
294 Ok((sched, submit_family, result_family))
295 }
296
297 fn build_transport_for(
302 family: crate::mmf_dispatcher::MmfFamily,
303 path: &Path,
304 capacity: usize,
305 ) -> Result<Arc<dyn MessageTransport>, SchedError> {
306 use crate::mmf_dispatcher::MmfFamily;
307 use crate::shared_deque::SharedDeque;
308 use crate::message_transport::PassSlot;
309 match family {
310 MmfFamily::SharedRing => {
311 let ring = SharedRing::open(path, capacity)
312 .or_else(|_| SharedRing::create(path, capacity))?;
313 Ok(Arc::new(ring))
314 }
315 MmfFamily::SharedDeque(_) => {
316 let deque = SharedDeque::<PassSlot>::create(path, capacity)
317 .map_err(|_| {
318 SchedError::UnsupportedTransportFamily(family)
319 })?;
320 Ok(Arc::new(deque))
321 }
322 MmfFamily::SharedHashMap => {
323 Err(SchedError::UnsupportedTransportFamily(family))
324 }
325 }
326 }
327
328 pub fn submitter(&self) -> Submitter {
330 Submitter::new(self.submit_ring.clone())
331 }
332
333 pub fn collector(&self) -> ResultCollector {
335 ResultCollector::new(self.result_ring.clone())
336 }
337
338 pub fn heartbeat(&self) -> Arc<HeartbeatTable> { self.heartbeat.clone() }
340
341 pub fn slot_idx(&self) -> usize { self.slot_idx }
343
344 pub fn watchdog_scan(&self) -> crate::failover::ReclaimReport {
347 let r = FailoverWatchdog::new(&self.heartbeat).scan();
348 self.ring_sidecar.push_op(
349 crate::sidecar_ops::scheduler::OP_WATCHDOG_SCAN,
350 if !r.is_empty() { 1 } else { 0 }, );
352 r
353 }
354
355 fn worker_loop(
356 submit: Arc<dyn MessageTransport>,
357 result: Arc<dyn MessageTransport>,
358 heartbeat: Arc<HeartbeatTable>,
359 slot_idx: usize,
360 stop: Arc<AtomicBool>,
361 ) {
362 let mut slot_buf = [0u8; PAYLOAD_BYTES];
363 let mut beat_counter = 0u64;
364 while !stop.load(Ordering::Acquire) {
365 beat_counter += 1;
366 if beat_counter.is_multiple_of(100) {
367 heartbeat.beat(slot_idx);
368 }
369 match submit.try_pop(&mut slot_buf) {
370 Ok(_) => {
371 let (pass, token) = parse_submit_slot(&slot_buf);
372 let bit = (token & 0x3F) as u8;
375 heartbeat.mark_in_flight(slot_idx, bit);
376 let r = exec_pass(&pass);
377 let result_payload = encode_result_slot(token, &r);
378 result.try_push(&result_payload).ok();
381 heartbeat.clear_in_flight(slot_idx, bit);
382 }
383 Err(TransportError::Empty) => {
384 std::thread::sleep(Duration::from_micros(100));
386 }
387 Err(_) => break,
388 }
389 }
390 heartbeat.unregister(slot_idx);
391 }
392}
393
394impl Drop for BackgroundScheduler {
395 fn drop(&mut self) {
396 self.stop.store(true, Ordering::Release);
397 if let Some(w) = self.worker.take() {
398 w.join().ok();
400 }
401 }
402}
403
404fn parse_submit_slot(slot: &[u8; PAYLOAD_BYTES]) -> (Pass, u32) {
405 let closure_id = u32::from_le_bytes(slot[0..4].try_into().unwrap());
406 let token = u32::from_le_bytes(slot[RESULT_TOKEN_OFFSET..ARG_LEN_OFFSET].try_into().unwrap());
407 let arg_len = u16::from_le_bytes(
408 slot[ARG_LEN_OFFSET..ARGS_OFFSET].try_into().unwrap()
409 ) as usize;
410 let args = slot[ARGS_OFFSET..ARGS_OFFSET + arg_len.min(MAX_ARG_LEN)].to_vec();
411 (Pass { closure_id, args }, token)
412}
413
414fn encode_result_slot(token: u32, r: &PassResult) -> [u8; PAYLOAD_BYTES] {
415 let mut buf = [0u8; PAYLOAD_BYTES];
416 buf[RESULT_TOKEN_R_OFFSET..STATUS_OFFSET].copy_from_slice(&token.to_le_bytes());
417 match r {
418 Ok(data) => {
419 buf[STATUS_OFFSET] = 0;
420 let len = data.len().min(MAX_RESULT_LEN);
421 buf[RESULT_LEN_OFFSET..RESULT_DATA_OFFSET]
422 .copy_from_slice(&(len as u16).to_le_bytes());
423 buf[RESULT_DATA_OFFSET..RESULT_DATA_OFFSET + len]
424 .copy_from_slice(&data[..len]);
425 }
426 Err(e) => {
427 buf[STATUS_OFFSET] = 1;
428 let msg = format!("{e:?}");
429 let bytes = msg.as_bytes();
430 let len = bytes.len().min(MAX_RESULT_LEN);
431 buf[RESULT_LEN_OFFSET..RESULT_DATA_OFFSET]
432 .copy_from_slice(&(len as u16).to_le_bytes());
433 buf[RESULT_DATA_OFFSET..RESULT_DATA_OFFSET + len]
434 .copy_from_slice(&bytes[..len]);
435 }
436 }
437 buf
438}
439
440#[cfg(test)]
441mod tests {
442 use super::*;
443 use crate::pass_registry;
444
445 fn tmp(name: &str) -> (std::path::PathBuf, std::path::PathBuf, std::path::PathBuf) {
446 let mut s = std::env::temp_dir(); let pid = std::process::id();
447 s.push(format!("subetha-sched-{name}-{pid}-submit.bin"));
448 let mut r = std::env::temp_dir();
449 r.push(format!("subetha-sched-{name}-{pid}-result.bin"));
450 let mut h = std::env::temp_dir();
451 h.push(format!("subetha-sched-{name}-{pid}-hb.bin"));
452 (s, r, h)
453 }
454
455 fn cleanup(paths: &(std::path::PathBuf, std::path::PathBuf, std::path::PathBuf)) {
456 std::fs::remove_file(&paths.0).ok();
457 std::fs::remove_file(&paths.1).ok();
458 std::fs::remove_file(&paths.2).ok();
459 }
460
461 #[test]
462 fn deque_backed_scheduler_round_trips_pass_end_to_end() {
463 use crate::message_transport::PassSlot;
467 use crate::shared_deque::SharedDeque;
468
469 let id = 0x2100_0002;
470 pass_registry::register(id, |args| {
471 Ok(args.iter().map(|b| b.wrapping_add(1)).collect())
472 });
473 let paths = tmp("deque_rt");
474 let submit_deque: Arc<dyn MessageTransport> = Arc::new(
475 SharedDeque::<PassSlot>::create(&paths.0, 64).expect("submit create"),
476 );
477 let result_deque: Arc<dyn MessageTransport> = Arc::new(
478 SharedDeque::<PassSlot>::create(&paths.1, 64).expect("result create"),
479 );
480 let sched = BackgroundScheduler::start_with_transports(
481 submit_deque,
482 result_deque,
483 &paths.2,
484 8,
485 )
486 .unwrap();
487 let submitter = sched.submitter();
488 let collector = sched.collector();
489 let token = submitter
490 .submit(&Pass {
491 closure_id: id,
492 args: vec![10, 20, 30],
493 })
494 .unwrap();
495 let mut got = None;
496 for _ in 0..200 {
497 if let Ok(r) = collector.try_recv() {
498 got = Some(r);
499 break;
500 }
501 std::thread::sleep(Duration::from_millis(2));
502 }
503 let r = got.expect("result must arrive within 400ms");
504 assert_eq!(r.token, token);
505 match r.result {
506 Ok(data) => assert_eq!(data, vec![11, 21, 31]),
507 Err(e) => panic!("expected Ok, got {e:?}"),
508 }
509 pass_registry::unregister(id);
510 drop(sched);
511 cleanup(&paths);
512 }
513
514 #[test]
515 fn workload_shape_routed_scheduler_round_trips_pass_end_to_end() {
516 use crate::dispatch_deque::WorkloadShape;
521 use crate::mmf_dispatcher::{MmfFamily, MmfWorkloadShape};
522
523 let id = 0x2100_0003;
524 pass_registry::register(id, |args| {
525 Ok(args.iter().map(|b| b.wrapping_mul(3)).collect())
526 });
527 let paths = tmp("by_shape_rt");
528 let submit_shape = MmfWorkloadShape::StreamingMpmc {
529 n_producers: 1,
530 n_consumers: 1,
531 };
532 let result_shape =
533 MmfWorkloadShape::WorkStealing(WorkloadShape::producer_fast(8));
534 let (sched, submit_family, result_family) =
535 BackgroundScheduler::start_by_workload_shape(
536 &paths.0,
537 submit_shape,
538 &paths.1,
539 result_shape,
540 &paths.2,
541 64,
542 8,
543 )
544 .expect("by_workload_shape");
545 assert_eq!(submit_family, MmfFamily::SharedRing);
547 assert!(matches!(result_family, MmfFamily::SharedDeque(_)));
549
550 let submitter = sched.submitter();
551 let collector = sched.collector();
552 let token = submitter
553 .submit(&Pass {
554 closure_id: id,
555 args: vec![5, 6, 7],
556 })
557 .unwrap();
558 let mut got = None;
559 for _ in 0..200 {
560 if let Ok(r) = collector.try_recv() {
561 got = Some(r);
562 break;
563 }
564 std::thread::sleep(Duration::from_millis(2));
565 }
566 let r = got.expect("result must arrive within 400ms");
567 assert_eq!(r.token, token);
568 match r.result {
569 Ok(data) => assert_eq!(data, vec![15, 18, 21]),
570 Err(e) => panic!("expected Ok, got {e:?}"),
571 }
572 pass_registry::unregister(id);
573 drop(sched);
574 cleanup(&paths);
575 }
576
577 #[test]
578 fn workload_shape_routed_scheduler_rejects_key_value_family() {
579 use crate::mmf_dispatcher::{MmfFamily, MmfWorkloadShape};
583
584 let paths = tmp("by_shape_reject_kv");
585 let bad_shape = MmfWorkloadShape::KeyValueLookup {
586 n_readers: 1,
587 n_writers: 1,
588 };
589 let good_shape = MmfWorkloadShape::StreamingMpmc {
590 n_producers: 1,
591 n_consumers: 1,
592 };
593 let result = BackgroundScheduler::start_by_workload_shape(
594 &paths.0,
595 bad_shape,
596 &paths.1,
597 good_shape,
598 &paths.2,
599 64,
600 8,
601 );
602 match result {
603 Err(SchedError::UnsupportedTransportFamily(MmfFamily::SharedHashMap)) => {
604 }
606 Err(other) => panic!("expected UnsupportedTransportFamily(SharedHashMap), got {other:?}"),
607 Ok(_) => panic!("expected error, got Ok"),
608 }
609 cleanup(&paths);
610 }
611
612 #[test]
613 fn submit_execute_result_round_trip() {
614 let id = 0x2000_0001;
615 pass_registry::register(id, |args| {
616 Ok(args.iter().map(|b| b.wrapping_mul(2)).collect())
618 });
619 let paths = tmp("rt");
620 let sched = BackgroundScheduler::start(
621 &paths.0, &paths.1, &paths.2, 64, 8,
622 ).unwrap();
623 let submitter = sched.submitter();
624 let collector = sched.collector();
625 let token = submitter.submit(&Pass {
626 closure_id: id, args: vec![1, 2, 3, 4],
627 }).unwrap();
628 let mut got = None;
630 for _ in 0..200 {
631 if let Ok(r) = collector.try_recv() {
632 got = Some(r);
633 break;
634 }
635 std::thread::sleep(Duration::from_millis(2));
636 }
637 let r = got.expect("result must arrive within 400ms");
638 assert_eq!(r.token, token);
639 match r.result {
640 Ok(data) => assert_eq!(data, vec![2, 4, 6, 8]),
641 Err(e) => panic!("expected Ok, got {e:?}"),
642 }
643 pass_registry::unregister(id);
644 drop(sched);
645 cleanup(&paths);
646 }
647
648 #[test]
649 fn submit_unknown_closure_returns_error() {
650 let paths = tmp("unknown");
651 let sched = BackgroundScheduler::start(
652 &paths.0, &paths.1, &paths.2, 64, 8,
653 ).unwrap();
654 let submitter = sched.submitter();
655 let collector = sched.collector();
656 let token = submitter.submit(&Pass {
657 closure_id: 0x9999_FFFF, args: vec![],
658 }).unwrap();
659 let mut got = None;
660 for _ in 0..200 {
661 if let Ok(r) = collector.try_recv() {
662 got = Some(r);
663 break;
664 }
665 std::thread::sleep(Duration::from_millis(2));
666 }
667 let r = got.expect("error result must arrive");
668 assert_eq!(r.token, token);
669 assert!(r.result.is_err());
670 drop(sched);
671 cleanup(&paths);
672 }
673
674 #[test]
675 fn args_too_large_rejected_at_submit() {
676 let paths = tmp("oversized");
677 let sched = BackgroundScheduler::start(
678 &paths.0, &paths.1, &paths.2, 8, 4,
679 ).unwrap();
680 let submitter = sched.submitter();
681 let big = vec![0u8; MAX_ARG_LEN + 1];
682 assert_eq!(
683 submitter.submit(&Pass { closure_id: 1, args: big }).unwrap_err(),
684 SchedError::ArgsTooLarge,
685 );
686 drop(sched);
687 cleanup(&paths);
688 }
689
690 #[test]
691 fn scheduler_registers_in_heartbeat_table() {
692 let paths = tmp("hb-reg");
693 let sched = BackgroundScheduler::start(
694 &paths.0, &paths.1, &paths.2, 8, 4,
695 ).unwrap();
696 let snap = sched.heartbeat().snapshot(sched.slot_idx()).unwrap();
697 assert_eq!(snap.pid, std::process::id());
698 drop(sched);
699 cleanup(&paths);
700 }
701
702 #[test]
703 fn many_passes_round_trip_in_order_of_submission() {
704 let id = 0x2000_0002;
705 pass_registry::register(id, |args| Ok(args.to_vec()));
706 let paths = tmp("many");
707 let sched = BackgroundScheduler::start(
708 &paths.0, &paths.1, &paths.2, 64, 4,
709 ).unwrap();
710 let submitter = sched.submitter();
711 let collector = sched.collector();
712 let mut tokens = Vec::new();
713 for i in 0..16u8 {
714 let t = submitter.submit(&Pass {
715 closure_id: id, args: vec![i; 4],
716 }).unwrap();
717 tokens.push(t);
718 }
719 let mut got = std::collections::HashMap::new();
720 let deadline = std::time::Instant::now() + Duration::from_secs(2);
721 while got.len() < tokens.len() && std::time::Instant::now() < deadline {
722 if let Ok(r) = collector.try_recv() {
723 got.insert(r.token, r.result);
724 } else {
725 std::thread::sleep(Duration::from_millis(2));
726 }
727 }
728 assert_eq!(got.len(), tokens.len());
729 for (i, t) in tokens.iter().enumerate() {
730 let r = got.get(t).expect("result for every token");
731 let data = r.as_ref().expect("Ok result");
732 assert_eq!(data, &vec![i as u8; 4]);
733 }
734 pass_registry::unregister(id);
735 drop(sched);
736 cleanup(&paths);
737 }
738}