Skip to main content

strat9_kernel/process/sched_classes/
fair.rs

1// SPDX-License-Identifier: MPL-2.0
2
3use super::{CurrentRuntime, SchedClassRq};
4use crate::process::task::Task;
5use alloc::{collections::BTreeMap, sync::Arc};
6use core::sync::atomic::Ordering;
7
8const WEIGHT_0: u64 = 1024;
9
10/// Base time slice per task in ticks for the CFS fair scheduler.
11///
12/// At TIMER_HZ=100 (10 ms/tick):
13///   BASE_SLICE_TICKS = 1 -> 1 tick = 10 ms per task (matches `quantum_ms: 10`)
14const BASE_SLICE_TICKS: u64 = 1;
15
16/// Fair starvation threshold in ticks.
17///
18/// A Fair task that has waited this many ticks without being selected is
19/// boosted to the front of its vruntime position.  At TIMER_HZ=100:
20/// 100 ticks = 1 second.  This prevents indefinite starvation when RT
21/// tasks consume most of the CPU.
22const FAIR_STARVATION_THRESHOLD_TICKS: u64 = 100;
23
24/// Performs the nice to weight operation.
25pub const fn nice_to_weight(nice: super::nice::Nice) -> u64 {
26    const FACTOR_NUMERATOR: u64 = 5;
27    const FACTOR_DENOMINATOR: u64 = 4;
28
29    const NICE_TO_WEIGHT: [u64; 40] = const {
30        let mut ret = [0; 40];
31        let mut index = 0;
32        let mut nice = super::nice::NiceValue::MIN.get();
33        while nice <= super::nice::NiceValue::MAX.get() {
34            ret[index] = match nice {
35                0 => WEIGHT_0,
36                nice @ 1.. => {
37                    let numerator = FACTOR_DENOMINATOR.pow(nice as u32);
38                    let denominator = FACTOR_NUMERATOR.pow(nice as u32);
39                    WEIGHT_0 * numerator / denominator
40                }
41                nice => {
42                    let numerator = FACTOR_NUMERATOR.pow((-nice) as u32);
43                    let denominator = FACTOR_DENOMINATOR.pow((-nice) as u32);
44                    WEIGHT_0 * numerator / denominator
45                }
46            };
47            index += 1;
48            nice += 1;
49        }
50        ret
51    };
52
53    NICE_TO_WEIGHT[(nice.value().get() + 20) as usize]
54}
55
56/// Per-CPU run queue for the Completely Fair Scheduler.
57///
58/// Uses two `BTreeMap`s for O(log n) operations without allocation on the
59/// fast paths (`pick_next`, `remove`):
60///
61/// - `entities`: primary map keyed by `(vruntime, task_id)` : the minimum
62///   entry is always the next task to schedule.
63/// - `by_id`: reverse index mapping `task_id => primary key` : enables O(log n)
64///   removal by task ID in `remove()` without scanning `entities`.
65///
66/// This replaces the previous `BinaryHeap`-based design which used *lazy
67/// deletion* (O(n) `remove()` scan, phantom entries, generation counters).
68/// The `BTreeMap` approach gives:
69///
70/// | Operation   | Complexity | Allocates?                              |
71/// |-------------|------------|-----------------------------------------|
72/// | `enqueue`   | O(log n)   | yes : 2 BTreeMap nodes (wakeup path)    |
73/// | `pick_next` | O(log n)   | no  : removes 2 nodes                   |
74/// | `remove`    | O(log n)   | no  : removes 2 nodes                   |
75///
76/// No phantom entries means no generation counter, no `prune_stale_head()`,
77/// and no per-entry liveness checks.  The BTreeMap pair is the authoritative
78/// record of which tasks are currently on the run queue.
79pub struct FairClassRq {
80    /// Primary index: `(vruntime, task_id)` => `(Arc<Task>, weight)`.
81    /// Ordered so `pop_first()` yields the task with the smallest vruntime.
82    /// `task_id` is part of the key to ensure uniqueness when two tasks share
83    /// the same vruntime.
84    entities: BTreeMap<(u64, u64), (Arc<Task>, u64)>,
85    /// Reverse index: `task_id` => primary key.
86    /// Allows `remove(task_id)` to locate and delete the `entities` entry in
87    /// O(log n) without scanning the primary map.
88    by_id: BTreeMap<u64, (u64, u64)>,
89    min_vruntime: u64,
90    total_weight: u64,
91    runnable_count: usize,
92}
93
94impl FairClassRq {
95    /// Creates a new instance.
96    pub fn new() -> Self {
97        Self {
98            entities: BTreeMap::new(),
99            by_id: BTreeMap::new(),
100            min_vruntime: 0,
101            total_weight: 0,
102            runnable_count: 0,
103        }
104    }
105
106    /// Total scheduling period in ticks.
107    ///
108    /// `BASE_SLICE_TICKS * (nr_runnable + 1)` : each runnable task gets at
109    /// least one full `BASE_SLICE_TICKS` per round.  `+1` accounts for the
110    /// currently-running task that is not counted in `runnable_count`.
111    fn period(&self) -> u64 {
112        let count = (self.runnable_count + 1) as u64;
113        (BASE_SLICE_TICKS * count).max(BASE_SLICE_TICKS)
114    }
115
116    /// Virtual-time slice: the vruntime budget for the current task.
117    fn vtime_slice(&self) -> u64 {
118        self.period() / (self.runnable_count + 1) as u64
119    }
120
121    /// Wall-clock time slice scaled by `cur_weight` relative to total weight.
122    fn time_slice(&self, cur_weight: u64) -> u64 {
123        let denom = self.total_weight + cur_weight;
124        if denom == 0 {
125            return self.period();
126        }
127        self.period() * cur_weight / denom
128    }
129}
130
131impl SchedClassRq for FairClassRq {
132    /// Enqueues a task onto the run queue.
133    ///
134    /// Clamps `vruntime` to `min_vruntime` so waking tasks do not receive an
135    /// unfair head start over tasks that have been waiting.  O(log n);
136    /// allocates two BTreeMap nodes.
137    fn enqueue(&mut self, task: Arc<Task>) {
138        if let super::SchedPolicy::Fair(nice) = task.sched_policy() {
139            let task_id = task.id.as_u64();
140
141            if self.by_id.contains_key(&task_id) {
142                return;
143            }
144
145            let weight = nice_to_weight(nice);
146            let mut vruntime = task.vruntime();
147            if vruntime < self.min_vruntime {
148                vruntime = self.min_vruntime;
149            }
150            task.set_vruntime(vruntime);
151            task.fair_prepare_enqueue();
152            // Reset starvation counter: task just arrived, hasn't waited yet.
153            task.fair_wait_ticks.store(0, Ordering::Relaxed);
154
155            let key = (vruntime, task_id);
156            self.entities.insert(key, (task, weight));
157            self.by_id.insert(task_id, key);
158            self.total_weight += weight;
159            self.runnable_count += 1;
160        }
161    }
162
163    /// Returns the number of tasks currently on the run queue.
164    fn len(&self) -> usize {
165        self.runnable_count
166    }
167
168    /// Picks the next task to run: the one with the smallest vruntime,
169    /// unless a starved task exists (waited > threshold), in which case
170    /// the starved task is boosted ahead.
171    ///
172    /// O(n) starvation scan + O(log n) normal pick.  The scan is bounded
173    /// by the number of Fair tasks (typically small).
174    fn pick_next(&mut self) -> Option<Arc<Task>> {
175        // Check for starved tasks: any task with fair_wait_ticks >= threshold
176        // is boosted ahead of normal vruntime ordering.
177        let mut starved_key: Option<(u64, u64)> = None;
178        let mut starved_wait: u64 = 0;
179        for (&key, (task, _)) in self.entities.iter() {
180            let wait = task.fair_wait_ticks.load(Ordering::Relaxed);
181            if wait >= FAIR_STARVATION_THRESHOLD_TICKS && wait > starved_wait {
182                starved_key = Some(key);
183                starved_wait = wait;
184            }
185        }
186
187        let key = if let Some(sk) = starved_key {
188            sk
189        } else {
190            // Normal path: pick minimum vruntime.
191            let (&k, _) = self.entities.iter().next()?;
192            k
193        };
194
195        let (task, weight) = self.entities.remove(&key)?;
196        self.by_id.remove(&key.1);
197        task.fair_mark_dequeued();
198        task.fair_wait_ticks.store(0, Ordering::Relaxed);
199        self.total_weight = self.total_weight.saturating_sub(weight);
200        self.runnable_count = self.runnable_count.saturating_sub(1);
201        Some(task)
202    }
203
204    /// Updates the vruntime of the currently-running task and decides whether
205    /// it should be preempted.
206    ///
207    /// Returns `true` if the task has exhausted its time slice or its vruntime
208    /// has overtaken the leftmost task's vruntime by more than `vtime_slice`.
209    fn update_current(&mut self, rt: &CurrentRuntime, task: &Task, is_yield: bool) -> bool {
210        if is_yield {
211            return true;
212        }
213        if let super::SchedPolicy::Fair(nice) = task.sched_policy() {
214            let weight = nice_to_weight(nice);
215            let delta_vruntime = if weight == 0 {
216                0
217            } else {
218                rt.delta_ticks * WEIGHT_0 / weight
219            };
220            let vruntime = task.vruntime() + delta_vruntime;
221            task.set_vruntime(vruntime);
222
223            // The leftmost entry is O(log n) to peek on a BTreeMap.
224            let leftmost_vruntime = self.entities.keys().next().map(|&(v, _)| v);
225            self.min_vruntime = match leftmost_vruntime {
226                Some(lv) => vruntime.min(lv),
227                None => vruntime,
228            };
229
230            // No other runnable task : keep running.
231            if leftmost_vruntime.is_none() {
232                return false;
233            }
234
235            rt.period_delta_ticks > self.time_slice(weight)
236                || vruntime > self.min_vruntime + self.vtime_slice()
237        } else {
238            false
239        }
240    }
241
242    /// Removes the task identified by `task_id` from the run queue.
243    ///
244    /// Uses the `by_id` reverse index for O(log n) lookup, then removes both
245    /// entries.  Allocation-free.  Returns `true` if the task was present.
246    fn remove(&mut self, task_id: crate::process::TaskId) -> bool {
247        let Some(key) = self.by_id.remove(&task_id.as_u64()) else {
248            return false;
249        };
250        if let Some((task, weight)) = self.entities.remove(&key) {
251            task.fair_invalidate_rq_entry();
252            self.total_weight = self.total_weight.saturating_sub(weight);
253            self.runnable_count = self.runnable_count.saturating_sub(1);
254            true
255        } else {
256            debug_assert!(
257                false,
258                "FairClassRq: by_id/entities out of sync for task {:?}",
259                task_id
260            );
261            false
262        }
263    }
264
265    /// Increment `fair_wait_ticks` for every queued task.  Called once per
266    /// timer tick from the LOCAL lock handler.
267    fn tick_update_wait(&mut self) {
268        for ((_, task_id), (task, _)) in self.entities.iter() {
269            let _ = task_id; // suppress unused warning
270            task.fair_wait_ticks.fetch_add(1, Ordering::Relaxed);
271        }
272    }
273}