Skip to main content

moirai_executor/schedule/route/
router.rs

1//! Static scheduler router.
2
3use core::marker::PhantomData;
4
5use moirai_core::Priority;
6
7use super::{
8    AcceleratorCounts, AcceleratorId, AcceleratorKind, AcceleratorRoute, AsyncLaneId, ProcessId,
9    ProcessRoute, RoutePolicy, RouteSummary, RouteTopology, SchedulerRoute, ServerId, ServerRoute,
10    ThreadId, ThreadRoute,
11};
12use crate::schedule::WorkClass;
13
14/// Static hybrid scheduler router.
15#[derive(Debug, Clone, Copy)]
16pub struct HybridRouter<P: RoutePolicy> {
17    topology: RouteTopology,
18    _policy: PhantomData<P>,
19}
20
21impl<P: RoutePolicy> HybridRouter<P> {
22    /// Construct a router for a topology and compile-time policy.
23    #[inline]
24    pub const fn new(topology: RouteTopology) -> Self {
25        Self {
26            topology,
27            _policy: PhantomData,
28        }
29    }
30
31    /// Return the router topology.
32    #[inline]
33    pub const fn topology(self) -> RouteTopology {
34        self.topology
35    }
36
37    /// Select one concrete scheduler route.
38    #[inline]
39    pub fn select<C: WorkClass>(&self, priority: Priority, sequence: usize) -> SchedulerRoute {
40        let topology = self.topology;
41        let priority = priority.index();
42        let route_key = sequence
43            .wrapping_add(C::AFFINITY_OFFSET)
44            .wrapping_add(priority);
45        let thread = ThreadId::new(route_key % topology.worker_threads().get());
46        let process = ProcessId::new(
47            sequence
48                .wrapping_div(topology.worker_threads().get())
49                .wrapping_add(C::AFFINITY_OFFSET)
50                .wrapping_add(priority)
51                % topology.processes().get(),
52        );
53        let async_lane = async_lane::<C>(topology, sequence, process);
54
55        if P::ENABLE_ACCELERATOR_ROUTES
56            && topology.accelerators().total() != 0
57            && route_key % P::ACCELERATOR_PERIOD.max(1) == 0
58        {
59            let accelerator_sequence = route_key / P::ACCELERATOR_PERIOD.max(1);
60            let (kind, accelerator) =
61                accelerator_target(topology.accelerators(), accelerator_sequence);
62            return SchedulerRoute::Accelerator(AcceleratorRoute {
63                kind,
64                accelerator,
65                process,
66                thread,
67                async_lane,
68            });
69        }
70
71        if P::ENABLE_SERVER_ROUTES
72            && topology.servers().get() != 0
73            && route_key % P::SERVER_PERIOD.max(1) == 0
74        {
75            return SchedulerRoute::Server(ServerRoute {
76                server: ServerId::new(route_key % topology.servers().get()),
77                process,
78                thread,
79                async_lane,
80            });
81        }
82
83        if P::ENABLE_PROCESS_ROUTES
84            && topology.processes().get() > 1
85            && (C::USES_ASYNC_LANE || route_key % P::PROCESS_PERIOD.max(1) == 0)
86        {
87            return SchedulerRoute::Process(ProcessRoute {
88                process,
89                thread,
90                async_lane,
91            });
92        }
93
94        SchedulerRoute::Thread(ThreadRoute { process, thread })
95    }
96
97    /// Summarize a deterministic route sequence for one work class.
98    #[inline]
99    pub fn summarize<C: WorkClass>(&self, priority: Priority, count: usize) -> RouteSummary {
100        let mut summary = RouteSummary::default();
101        for sequence in 0..count {
102            summary.record(self.select::<C>(priority, sequence));
103        }
104        summary
105    }
106}
107
108#[inline]
109fn async_lane<C: WorkClass>(
110    topology: RouteTopology,
111    sequence: usize,
112    process: ProcessId,
113) -> Option<AsyncLaneId> {
114    if C::USES_ASYNC_LANE {
115        Some(AsyncLaneId::new(
116            sequence.wrapping_add(process.get()) % topology.async_lanes_per_process().get(),
117        ))
118    } else {
119        None
120    }
121}
122
123#[inline]
124fn accelerator_target(
125    accelerators: AcceleratorCounts,
126    accelerator_sequence: usize,
127) -> (AcceleratorKind, AcceleratorId) {
128    let target = accelerator_sequence % accelerators.total();
129    if target < accelerators.cpu() {
130        return (AcceleratorKind::Cpu, AcceleratorId::new(target));
131    }
132    let target = target - accelerators.cpu();
133    if target < accelerators.gpu() {
134        return (AcceleratorKind::Gpu, AcceleratorId::new(target));
135    }
136    let target = target - accelerators.gpu();
137    if target < accelerators.tpu() {
138        return (AcceleratorKind::Tpu, AcceleratorId::new(target));
139    }
140    (
141        AcceleratorKind::Npu,
142        AcceleratorId::new(target - accelerators.tpu()),
143    )
144}