Learn Rust Series (#71) - Custom Thread Pools and Work Stealing

Words
2881
Reading
13 min
Listen
Play
39m

Learn Rust Series (#71) - Custom Thread Pools and Work Stealing

rust-banner.png

What will I learn

  • You will learn why spawning one OS thread per task is wasteful, and what a thread pool fixes;
  • how to build a pool from a channel, a shared receiver, and a fixed set of worker threads;
  • how to submit closures as jobs with Box<dyn FnOnce> and collect their results;
  • how a pool shuts down cleanly by closing the channel and joining every worker;
  • what work stealing is, and how a shared queue approximates it in plain std.

Requirements

  • A working modern computer running macOS, Windows or Ubuntu;
  • An installed Rust toolchain (via rustup, from rustup.rs);
  • The previous seventy episodes, especially channels, Arc, Mutex, and scoped threads;
  • The ambition to learn systems programming from the ground up.

Difficulty

  • Intermediate

Curriculum (of the Learn Rust Series):

Learn Rust Series (#71) - Custom Thread Pools and Work Stealing

Creating a thread is not free. When you call thread::spawn, the operating system has to allocate a stack (often a couple of megabytes), register the thread with its scheduler, and wire up the machinery to run it. That is real time, and for one big long-running job it is neglible. But spin up one thread per tiny task -- one per incoming request, one per pixel -- and the setup cost swamps the useful work. A thread pool fixes this by starting a fixed set of worker threads exactly once, then feeding them a stream of jobs through a channel. It is the pattern behind every web server, every parallel runtime, and (as we saw last episode) the machinery hiding inside rayon. Building one by hand ties together channels from episode 63, Arc and Mutex from episodes 34 and 65, and the boxed closures we have used since episode 11, into one satisfying whole ;-)

Solutions to Episode 70 Exercises

Episode 70 was rayon and data parallelism. Here are worked solutions to the three exercises.

Exercise 1 -- a parallel_sum that splits a slice across scoped threads using ceil-division:

use std::thread;

fn parallel_sum(data: &[i64], threads: usize) -> i64 {
    let chunk = (data.len() + threads - 1) / threads; // ceil division
    thread::scope(|s| {
        data.chunks(chunk)
            .map(|c| s.spawn(move || c.iter().sum::<i64>()))
            .collect::<Vec<_>>()
            .into_iter()
            .map(|h| h.join().unwrap())
            .sum()
    })
}

fn main() {
    let data: Vec<i64> = (1..=100).collect();
    println!("{}", parallel_sum(&data, 4)); // 5050
}

The ceil-division is the whole trick. Plain integer division rounds the chunk size down, and the leftover tail elements would silently never be summed. Rounding up guarantees every element lands in some chunk, even when the length does not divide evenly.

Exercise 2 -- a recursive divide-and-conquer parallel max with a base case and the "run one half on the current thread" trick:

use std::thread;

fn par_max(data: &[i32]) -> Option<i32> {
    if data.len() <= 128 {
        return data.iter().copied().max(); // base case: too small to split
    }
    let mid = data.len() / 2;
    let (l, r) = data.split_at(mid);
    thread::scope(|s| {
        let left = s.spawn(|| par_max(l)); // spawn one half
        let right = par_max(r);            // run the other on THIS thread
        [left.join().unwrap(), right].into_iter().flatten().max()
    })
}

fn main() {
    let data: Vec<i32> = (0..1000).rev().collect();
    println!("{:?}", par_max(&data)); // Some(999)
}

Two details carry the correctness. The <= 128 base case stops us spawning a thread for a two-element slice, where the coordination would cost more than the comparison. And running the right half on the current thread, rather than spawning a second worker, keeps the caller busy instead of blocking it idle -- the exact trick rayon's join uses internally.

Exercise 3 -- describing where par_iter helps and where it does not:

// HELPS: rendering a 4000x4000 Mandelbrot fractal. Each pixel runs an
// independent iteration loop of potentially hundreds of steps, so per-element
// work is large and the collection (16 million pixels) is huge. The split cost
// is nothing next to the compute, and every core stays saturated.
//
// DOES NOT HELP: summing a Vec of a few thousand bytes. Per-element work is
// a single add, the whole slice fits in cache, and the job is memory-bound.
// Splitting, waking threads, and combining costs MORE than the sum itself, so
// par_iter loses to a plain iter().sum().

The single question to ask is whether the per-element work dwarfs the coordination overhead. If it does, and the collection is large, parallelism wins. If the work is trivial or the data is small, the split costs more than it saves. Now, thread pools.

The cost a pool avoids

Spawning a thread per task works, and for a handful of large jobs it is completely fine. The problem is scale:

use std::thread;

fn main() {
    // fine for four big jobs, wasteful for four thousand tiny ones:
    let handles: Vec<_> = (0..4).map(|i| thread::spawn(move || i * 2)).collect();
    let doubled: Vec<i32> = handles.into_iter().map(|h| h.join().unwrap()).collect();
    println!("{doubled:?}"); // [0, 2, 4, 6]
}

For four jobs this is perfect. For four thousand tiny jobs it becomes a disaster: you pay the stack-allocation and scheduler-registration cost four thousand times, and you may have four thousand threads alive at once, each demanding memory and a scheduling slot. A pool inverts this. It pays the thread-creation cost N times up front (once per worker), then reuses those same N threads across as many jobs as you throw at it. The jobs become cheap messages; the expensive threads are created once and live for the whole run.

A thread pool from scratch

The design has four moving parts. A channel carries jobs. The Sender half lives in the pool; the Receiver half is shared among all the workers, wrapped in Arc<Mutex<...>> so many threads can hold it but only one pulls a job at a time. Each worker is a thread looping forever: lock the receiver, take the next job, release the lock, run it. And a job is just a boxed closure, Box<dyn FnOnce() + Send + 'static>, because we do not know at compile time what work the caller will hand us:

use std::sync::{mpsc, Arc, Mutex};
use std::thread;

type Job = Box<dyn FnOnce() + Send + 'static>;

struct ThreadPool {
    workers: Vec<thread::JoinHandle>,
    sender: Option<mpsc::Sender>,
}

impl ThreadPool {
    fn new(size: usize) -> Self {
        let (sender, receiver) = mpsc::channel::();
        let receiver = Arc::new(Mutex::new(receiver));
        let mut workers = Vec::with_capacity(size);
        for _ in 0..size {
            let receiver = Arc::clone(&receiver);
            workers.push(thread::spawn(move || {
                while let Ok(job) = receiver.lock().unwrap().recv() {
                    job(); // the lock is released before the job runs
                }
            }));
        }
        ThreadPool { workers, sender: Some(sender) }
    }

    fn execute<F: FnOnce() + Send + 'static>(&self, f: F) {
        self.sender.as_ref().unwrap().send(Box::new(f)).unwrap();
    }
}

impl Drop for ThreadPool {
    fn drop(&mut self) {
        drop(self.sender.take()); // close the channel so workers leave their loop
        for w in self.workers.drain(..) {
            w.join().unwrap();
        }
    }
}

fn main() {
    let pool = ThreadPool::new(4);
    for i in 0..8 {
        pool.execute(move || println!("task {i} ran"));
    }
    // pool dropped here: pending jobs finish, then every worker joins
}

A few subtleties are worth slowing down on. First, job() runs after the lock guard is dropped. Look closely: receiver.lock().unwrap().recv() produces a temporary MutexGuard, and because we never bind it to a variable, it is dropped at the end of the while let condition -- before the loop body runs. So a worker holds the lock only long enough to grab a job, not while executing it, which leaves the other workers free to pull their own jobs in parallel. Get this wrong -- hold the guard across job() -- and your "pool" runs everything serially, one job at a time, which defeats the whole point.

Second, the Job type alias hides three requirements the compiler will hold you to: FnOnce (the closure runs once and is consumed), Send (it can cross a thread boundary), and 'static (it holds no borrowed references that could dangle). A closure that captures a non-Send value simply will not compile -- the fearless-concurrency guarantee from episode 61 doing its job so you do not have to think about it.

Clean shutdown without hanging

The Option<Sender> wrapper and the Drop implementation are not decoration -- together they prevent a real deadlock. If the pool held a plain Sender and we dropped the pool, we would try to join the workers while the channel was still open. But an mpsc channel only closes once all senders are dropped, and the pool still owns one, so every worker's recv() would block forever and join would hang the program. By storing the sender in an Option and calling take() first, we drop it explicitly, which closes the channel. Each worker's recv() then returns Err, the while let loops end, and join returns cleanly.

Notice the ordering inside drop: we close the channel before joining. Reverse it -- join first, then close -- and you are back to the hang, because the workers are still waiting on an open channel when you start waiting on them. This is exactly the kind of ordering bug that episode 42 (Drop order and leak safety) trained your eye to catch.

Getting results back

The pool above runs side effects: printing, writing, mutating shared state. To collect computed values back out, add a second channel the jobs send their results into, and drain it on the main thread:

use std::sync::mpsc;
use std::thread;

fn main() {
    let (result_tx, result_rx) = mpsc::channel();
    let mut handles = Vec::new();
    for i in 1..=4 {
        let tx = result_tx.clone();
        handles.push(thread::spawn(move || tx.send(i * i).unwrap()));
    }
    drop(result_tx);
    let mut results: Vec<i32> = result_rx.iter().collect();
    for h in handles { h.join().unwrap(); }
    results.sort();
    println!("{results:?}"); // [1, 4, 9, 16]
}

The drop(result_tx) before draining is the same closing trick again. result_rx.iter() yields until every sender is gone, so if we forgot to drop our own copy of the sender, the iterator would block forever waiting for a result that never comes. Concurrency in Rust is full of these "who still holds a sender?" questions, and once the pattern clicks it becomes second nature.

Watching the load spread

One quietly wonderful property of the shared-queue design is that it load-balances for free. Because every idle worker races to lock the receiver and grab the next job, a worker that finishes a quick job comes straight back for more, while a worker stuck on a slow job simply does not. Fast workers naturally handle more jobs; slow workers handle fewer. Nobody schedules this -- it falls out of the shared queue. Here three workers pull thirty jobs and we confirm every one was handled:

use std::sync::{mpsc, Arc, Mutex};
use std::thread;

fn main() {
    let (tx, rx) = mpsc::channel::<i32>();
    let rx = Arc::new(Mutex::new(rx));
    let counts = Arc::new(Mutex::new([0usize; 3])); // per-worker job counts
    thread::scope(|s| {
        for id in 0..3 {
            let (rx, counts) = (Arc::clone(&rx), Arc::clone(&counts));
            s.spawn(move || {
                while rx.lock().unwrap().recv().is_ok() {
                    counts.lock().unwrap()[id] += 1;
                }
            });
        }
        for j in 0..30 { tx.send(j).unwrap(); }
        drop(tx);
    });
    let total: usize = counts.lock().unwrap().iter().sum();
    println!("all {total} jobs handled across the pool"); // 30
}

If you print the per-worker counts you will rarely see a clean [10, 10, 10]. On a real machine, with the OS scheduler and cache effects in play, you get something lumpy like [12, 9, 9] -- and that lumpiness is the load balancing working, not failing. The workers that happened to be free took more.

From shared queue to work stealing

The single shared queue has one real weakness: contention. Every worker fights for the same mutex on every single job. With four workers that is tolerable; with sixty-four it becomes a bottleneck, because the lock itself serializes all the workers at the exact moment they are trying to run in parallel. Work stealing is the industrial-strength fix. Instead of one shared queue, each worker gets its own seperate local double-ended queue. A worker pushes and pops from its own end with no locking in the common case; only when a worker runs dry does it reach over and steal a task from the back of a busy neighbour's queue. That is exactly how rayon (last episode) and Tokio's scheduler balance load without a central bottleneck, and the crossbeam-deque crate provides the primitives:

// requires the `crossbeam-deque` crate: shown for illustration, not compiled locally
use crossbeam_deque::{Stealer, Worker};

fn main() {
    let worker: Worker<i32> = Worker::new_fifo();
    worker.push(1);
    worker.push(2);
    let stealer: Stealer<i32> = worker.stealer(); // idle workers steal via this handle
    let _ = stealer;
    println!("{:?}", worker.pop()); // Some(1)
}

The Worker is the owner's end, cheap and lock-free for the owner; the Stealer is a cloneable handle other workers use to take from the far end. Because the owner works one end and thieves work the other, the two rarely collide, and when they do the deque resolves it with the compare-and-swap atomics from episode 66 rather than a lock.

You can approximate the idea in pure std with a single shared VecDeque that every worker pops from. It is simpler and still balances load; it just keeps the central lock that real work stealing eliminates:

use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use std::thread;

fn main() {
    let queue = Arc::new(Mutex::new((0..20).collect::<VecDeque<i32>>()));
    let done = Arc::new(Mutex::new(Vec::new()));
    thread::scope(|s| {
        for _ in 0..4 {
            let (q, d) = (Arc::clone(&queue), Arc::clone(&done));
            s.spawn(move || loop {
                let task = q.lock().unwrap().pop_front();
                match task {
                    Some(n) => d.lock().unwrap().push(n * n),
                    None => break, // the queue is drained
                }
            });
        }
    });
    println!("processed {} tasks", done.lock().unwrap().len()); // 20
}

Notice that we lock, pop, and release before doing the work (n * n) outside the lock -- the same discipline as the pool's worker loop. The step from here to true work stealing is really just "give each worker its own local deque and let idle ones raid the others," which is more code than an episode can hold but no new concepts.

How Python and Go would frame this

If you come from Python, the direct analogue is concurrent.futures, but with the enormous caveat of the global interpreter lock. For CPU-bound work the GIL means threads do not run Python bytecode in parallel, so you reach for a process pool instead:

from concurrent.futures import ProcessPoolExecutor

def square(n):
    return n * n

if __name__ == "__main__":
    with ProcessPoolExecutor(max_workers=4) as pool:
        results = list(pool.map(square, range(10)))
    print(results)  # [0, 1, 4, 9, 16, 25, 36, 49, 64, 81]

executor.submit is the moral equivalent of our pool.execute, and collecting the results is our results channel. But Python hands you the pool as a batteries-included object, whereas we built ours -- and in building it you saw every part the executor hides, plus the process-vs-thread tax that Rust never pays because it shares memory safely.

Go bakes the pool pattern right into the language's shape: a buffered channel of jobs, a fixed number of goroutines ranging over it, and a channel for results:

jobs := make(chan int, 100)
results := make(chan int, 100)

for w := 0; w < 4; w++ { // four workers
    go func() {
        for j := range jobs {
            results <- j * j
        }
    }()
}

for i := 1; i <= 10; i++ {
    jobs <- i
}
close(jobs) // like dropping our Sender: the workers' range loops end

for i := 0; i < 10; i++ {
    fmt.Println(<-results)
}

This is almost line for line our Rust design, which is no coincidence: it is the canonical concurrency pattern, and every language expresses the same shape. Go's version is beautifully terse. Rust asks for a little more ceremony -- Arc, Mutex, the Drop -- and in return the compiler proves you never shared anything unsafely, and the close(jobs)-equals-drop(sender) parallel shows the shutdown logic is identical underneath.

Build one, then reach for the crate

Here is the honest advice to close on. You will almost never write a thread pool in production Rust, and that is fine -- rayon gives you a work-stealing pool for data parallelism, tokio gives you one for async tasks, and both are battle-tested in ways a hand-rolled version never will be. So why build one at all? Because now, when you use rayon and watch it saturate your cores, or when a Tokio task mysteriously stalls the whole runtime, you know precisely what is happening under the surface: a fixed set of workers, a queue of jobs, a lock or a steal, a clean shutdown. You built the thing, so the thing holds no mystery. That is the entire reason we spend an episode reinventing a wheel the ecosystem already ships -- not to use our wheel, but to understand every wheel we will ever roll on. Next we look at what happens when a plain queue is not enough coordination, and workers need to wait for a condition or rendezvous at a shared point before continuing.

Bedankt en tot de volgende keer!

We opened with a simple observation -- threads are expensive, jobs are cheap -- and closed with a working pool, a deadlock-free shutdown, a results channel, and a clear line of sight from our shared-queue toy all the way up to rayon's work-stealing scheduler. If you built the pool yourself as you read (and I really hope you did), you now own a mental model that most people who merely use thread pools never acquire. Hou vol, blijf bouwen, en tot de volgende keer ;-)

Exercises

  1. Build a ThreadPool with a new(size) constructor and an execute method, then use it to run ten jobs that each print their number. Confirm the pool shuts down cleanly when it is dropped at the end of main.
  2. Extend your pool so each job sends its result into a channel; after submitting all jobs, drop your own sender and collect every result into a sorted Vec<i32>.
  3. In a comment, explain why holding the receiver's MutexGuard across the call to job() would turn your parallel pool into a sequential one.

scipio@scipio

Learn Rust Series (#71) - Custom Thread Pools and Work Stealing | Ecency