Skip to main content

ruda_runtime/runtime/stream/
scheduler.rs

1use crate::runtime::{
2    config::streaming::StreamingLogLevel,
3    logging::ServerLogger,
4    stream::{StreamFactory, StreamPool},
5};
6use alloc::{format, sync::Arc, vec::Vec};
7use ruda_core::stream_id::StreamId;
8
9/// Defines a trait for a scheduler stream backend, specifying the types and behavior for task scheduling.
10pub trait SchedulerStreamBackend {
11    /// Type representing a task.
12    type Task: core::fmt::Debug;
13    /// Type representing a stream.
14    type Stream: core::fmt::Debug;
15    /// Type for the stream factory, which creates streams of type `Self::Stream`.
16    type Factory: StreamFactory<Stream = Self::Stream>;
17
18    /// Enqueues a task onto a given stream for execution.
19    fn enqueue(task: Self::Task, stream: &mut Self::Stream);
20    /// Flush the inner stream queue to ensure ordering between different streams.
21    fn flush(stream: &mut Self::Stream);
22    /// Returns a mutable reference to the stream factory.
23    fn factory(&mut self) -> &mut Self::Factory;
24}
25
26/// Represents a multi-stream scheduler that manages task execution across multiple streams.
27#[derive(Debug)]
28pub struct SchedulerMultiStream<B: SchedulerStreamBackend> {
29    /// Pool of streams managed by the scheduler.
30    pool: StreamPool<SchedulerPoolMarker<B>>,
31    /// Strategy for scheduling tasks (e.g., Interleave or Sequential).
32    strategy: SchedulerStrategy,
33    /// Maximum number of tasks allowed per stream before execution is triggered.
34    max_tasks: usize,
35    /// Server logger.
36    pub logger: Arc<ServerLogger>,
37}
38
39/// Defines the scheduling strategy for task execution.
40#[derive(Debug)]
41pub enum SchedulerStrategy {
42    /// Tasks from different streams are interleaved during execution.
43    Interleave,
44    /// Tasks from each stream are executed sequentially.
45    Sequential,
46}
47
48/// Represents a single stream that holds tasks and a backend stream.
49#[derive(Debug)]
50pub struct Stream<B: SchedulerStreamBackend> {
51    /// List of tasks queued for execution in this stream.
52    tasks: Vec<B::Task>,
53    /// The backend stream used for task execution.
54    stream: B::Stream,
55}
56
57impl<B: SchedulerStreamBackend> Stream<B> {
58    /// Flushes all tasks from the stream, returning them and clearing the internal task list.
59    fn flush(&mut self) -> Vec<B::Task> {
60        if self.tasks.is_empty() {
61            return Vec::new();
62        }
63        let mut returned = Vec::with_capacity(self.tasks.capacity());
64        core::mem::swap(&mut returned, &mut self.tasks);
65        returned
66    }
67}
68
69#[derive(Debug)]
70struct SchedulerPoolMarker<B: SchedulerStreamBackend> {
71    backend: B,
72}
73
74impl<B: SchedulerStreamBackend> StreamFactory for SchedulerPoolMarker<B> {
75    // The type of stream produced by this factory.
76    type Stream = Stream<B>;
77
78    // Creates a new stream with an empty task list and a backend stream.
79    fn create(&mut self) -> Self::Stream {
80        Stream {
81            tasks: Vec::new(),
82            // Uses the backend's factory to create a new stream.
83            stream: self.backend.factory().create(),
84        }
85    }
86}
87
88/// Options for configuring a `SchedulerMultiStream`.
89#[derive(Debug)]
90pub struct SchedulerMultiStreamOptions {
91    /// Maximum number of streams allowed in the pool.
92    pub max_streams: u8,
93    /// Maximum number of tasks per stream before execution is triggered.
94    pub max_tasks: usize,
95    /// The scheduling strategy to use.
96    pub strategy: SchedulerStrategy,
97}
98
99impl<B: SchedulerStreamBackend> SchedulerMultiStream<B> {
100    /// Creates a new `SchedulerMultiStream` with the given backend and options.
101    pub fn new(
102        logger: Arc<ServerLogger>,
103        backend: B,
104        options: SchedulerMultiStreamOptions,
105    ) -> Self {
106        Self {
107            pool: StreamPool::new(SchedulerPoolMarker { backend }, options.max_streams, 0),
108            max_tasks: options.max_tasks,
109            strategy: options.strategy,
110            logger,
111        }
112    }
113
114    /// Returns a mutable reference to the backend stream for a given stream ID.
115    pub fn stream(&mut self, stream_id: &StreamId) -> &mut B::Stream {
116        let stream = self.pool.get_mut(stream_id);
117        &mut stream.stream
118    }
119
120    /// Registers a task for execution on a specific stream, ensuring stream alignment.
121    pub fn register(&mut self, stream_id: StreamId, task: B::Task, args_streams: &[StreamId]) {
122        // Align streams to ensure dependencies are handled correctly.
123        self.align_streams(stream_id, args_streams);
124
125        // Get the stream for the given stream ID and add the task to its queue.
126        let current = self.pool.get_mut(&stream_id);
127        current.tasks.push(task);
128
129        // If the task queue exceeds the maximum, execute the stream.
130        if current.tasks.len() >= self.max_tasks {
131            let index = self.pool.stream_index(&stream_id);
132            self.execute_stream_index(index);
133        }
134    }
135
136    /// Aligns streams by flushing tasks from streams that conflict with the given bindings.
137    pub(crate) fn align_streams(&mut self, stream_id: StreamId, args_streams: &[StreamId]) {
138        let mut to_flush = Vec::new();
139        // Get the index of the target stream.
140        let index = self.pool.stream_index(&stream_id);
141
142        // Identify streams that need to be flushed due to conflicting bindings.
143        for arg_stream in args_streams {
144            let index_stream = self.pool.stream_index(arg_stream);
145            if index != index_stream {
146                to_flush.push(*arg_stream);
147
148                self.logger.log_streaming(
149                    |level| matches!(level, StreamingLogLevel::Full),
150                    || format!("Binding on {} is shared on {}", arg_stream, stream_id),
151                );
152            }
153        }
154
155        // If no streams need flushing, return early.
156        if to_flush.is_empty() {
157            return;
158        }
159
160        self.logger.log_streaming(
161            |level| !matches!(level, StreamingLogLevel::Disabled),
162            || {
163                format!(
164                    "Flushing streams {to_flush:?} before registering more tasks on {stream_id}"
165                )
166            },
167        );
168        // Execute the streams that need to be flushed.
169        self.execute_streams(to_flush);
170    }
171
172    /// Executes tasks from the specified streams based on the scheduling strategy.
173    pub fn execute_streams(&mut self, stream_ids: Vec<StreamId>) {
174        if let [id] = stream_ids.as_slice() {
175            let index = self.pool.stream_index(id);
176            drop(stream_ids);
177            self.execute_stream_index(index);
178            return;
179        }
180        let mut indices = Vec::with_capacity(stream_ids.len());
181        let mut seen = [0u64; 4];
182
183        // Collect unique stream indices to avoid redundant processing.
184        for id in stream_ids {
185            let index = self.pool.stream_index(&id);
186            let word = &mut seen[index / 64];
187            let mask = 1u64 << (index % 64);
188            if *word & mask == 0 {
189                *word |= mask;
190                indices.push(index);
191            }
192        }
193
194        if let [index] = indices.as_slice() {
195            let index = *index;
196            drop(indices);
197            self.execute_stream_index(index);
198            return;
199        }
200
201        // Create schedules for each stream to be executed.
202        let mut schedules = Vec::new();
203        for index in indices {
204            let stream = unsafe { self.pool.get_mut_index(index) }; // Note: `unsafe` usage assumes valid index.
205            let tasks = stream.flush();
206            let num_tasks = tasks.len();
207
208            schedules.push(Schedule {
209                tasks: tasks.into_iter(),
210                num_tasks,
211                stream_index: index,
212            });
213        }
214
215        // If no schedules were created, return early.
216        if schedules.is_empty() {
217            return;
218        }
219
220        // Execute schedules based on the configured strategy.
221        match self.strategy {
222            SchedulerStrategy::Interleave => self.execute_schedules_interleave(schedules),
223            SchedulerStrategy::Sequential => self.execute_schedules_sequence(schedules),
224        }
225    }
226
227    fn execute_stream_index(&mut self, index: usize) {
228        let stream = unsafe { self.pool.get_mut_index(index) };
229        let tasks = stream.flush();
230        let schedule = Schedule {
231            num_tasks: tasks.len(),
232            tasks: tasks.into_iter(),
233            stream_index: index,
234        };
235        match self.strategy {
236            SchedulerStrategy::Interleave => self.execute_schedules_interleave([schedule]),
237            SchedulerStrategy::Sequential => self.execute_schedules_sequence([schedule]),
238        }
239    }
240
241    /// Executes schedules sequentially, processing each stream's tasks in order.
242    fn execute_schedules_sequence(&mut self, schedules: impl IntoIterator<Item = Schedule<B>>) {
243        for schedule in schedules {
244            let stream = unsafe { self.pool.get_mut_index(schedule.stream_index) }; // Note: `unsafe` usage assumes valid index.
245            for task in schedule.tasks {
246                // Enqueue each task on the stream.
247                B::enqueue(task, &mut stream.stream);
248            }
249
250            // Makes sure the tasks are ordered on the compute queue.
251            B::flush(&mut stream.stream);
252        }
253    }
254
255    //// Executes schedules in an interleaved manner, alternating tasks from different streams.
256    ///
257    /// We chose the first stream as the one executing the tasks, ensuring proper ordering by
258    /// flushing all other streams first and flushing the execution stream at the end.
259    /// This way, we ensure that most tasks are actually interleaved on the real compute queue
260    /// shared across all streams.
261    fn execute_schedules_interleave(&mut self, mut schedules: impl AsMut<[Schedule<B>]>) {
262        let schedules = schedules.as_mut();
263        // Makes sure the tasks are ordered on the compute queue.
264        for schedule in schedules.iter_mut().skip(1) {
265            let stream = unsafe { self.pool.get_mut_index(schedule.stream_index) };
266            B::flush(&mut stream.stream);
267        }
268
269        let execution_index = schedules.first().expect("At least one stream").stream_index;
270        let stream = unsafe { self.pool.get_mut_index(execution_index) };
271
272        // Find the maximum number of tasks across all schedules.
273        let num_tasks_max = schedules
274            .iter()
275            .map(|s| s.num_tasks)
276            .max()
277            .expect("At least one schedule");
278
279        // Iterate through tasks, interleaving them across streams.
280        for _ in 0..num_tasks_max {
281            for schedule in schedules.iter_mut() {
282                // If there are tasks remaining in the schedule, enqueue the next one.
283                if let Some(task) = schedule.tasks.next() {
284                    B::enqueue(task, &mut stream.stream);
285                }
286            }
287        }
288
289        // Making sure all tasks are registered to the queue.
290        B::flush(&mut stream.stream);
291    }
292}
293
294// Represents a schedule for executing tasks on a specific stream.
295struct Schedule<B: SchedulerStreamBackend> {
296    // Iterator over the tasks to be executed.
297    tasks: alloc::vec::IntoIter<B::Task>,
298    // Number of tasks in the schedule.
299    num_tasks: usize,
300    // Index of the stream in the pool.
301    stream_index: usize,
302}