Skip to main content

wowlab_node/worker/
pool.rs

1use std::sync::{
2    Arc, Mutex, MutexGuard, PoisonError,
3    atomic::{AtomicU32, AtomicU64, Ordering::SeqCst},
4};
5
6use tokio::{
7    runtime::Handle,
8    sync::{OwnedSemaphorePermit, Semaphore, mpsc},
9};
10use uuid::Uuid;
11use wowlab_common::{
12    ClaimToken, RuntimeChunkId, RuntimeWorkItem, RuntimeWorkResult, WorkContextHash, time::Instant,
13};
14use wowlab_engine_ports::{ContentCatalog, DataResolver, DynDataResolver};
15use wowlab_types::constants::SECONDS_PER_MINUTE;
16
17use super::runner::run_item;
18use crate::{NodeStats, work_context::WorkContext};
19
20const WORK_CHANNEL_SIZE: usize = 100;
21const RESULT_CHANNEL_SIZE: usize = 100;
22
23fn log_item_failure(job_id: Uuid, item_id: u64, error: &impl std::fmt::Display) {
24    tracing::error!(%job_id, item_id, %error, "Work item failed");
25}
26
27fn log_task_failure(job_id: Uuid, error: &impl std::fmt::Display) {
28    tracing::error!(%job_id, %error, "Work item task failed");
29}
30
31/// A resolved chunk of abstract work items queued for execution.
32#[derive(Debug)]
33pub struct WorkBatch {
34    pub job_id: Uuid,
35    pub chunk_id: RuntimeChunkId,
36    pub claim_token: ClaimToken,
37    pub work_context_hash: WorkContextHash,
38    pub context: WorkContext,
39    pub items: Vec<RuntimeWorkItem>,
40}
41
42/// Result of executing a whole chunk: one compact result per work item.
43#[derive(Debug)]
44pub struct WorkBatchResult {
45    pub job_id: Uuid,
46    pub chunk_id: RuntimeChunkId,
47    pub claim_token: ClaimToken,
48    pub work_context_hash: WorkContextHash,
49    pub results: Vec<RuntimeWorkResult>,
50    pub elapsed_ms: u64,
51}
52
53/// Thread pool that executes simulation chunks across available CPU cores.
54pub struct WorkerPool {
55    catalog: &'static ContentCatalog,
56    limiter: Arc<WorkerLimiter>,
57    completed_chunks: Arc<AtomicU64>,
58    sims_completed: Arc<AtomicU64>,
59    work_tx: Option<mpsc::Sender<WorkBatch>>,
60    result_rx: Option<mpsc::Receiver<WorkBatchResult>>,
61    resolver: Option<Arc<DynDataResolver<'static>>>,
62}
63
64#[derive(Debug)]
65struct WorkerLimitState {
66    limit: usize,
67    capacity: usize,
68}
69
70#[derive(Debug)]
71struct WorkerLimiter {
72    semaphore: Arc<Semaphore>,
73    state: Mutex<WorkerLimitState>,
74    active: AtomicU32,
75}
76
77impl WorkerLimiter {
78    fn new(limit: usize) -> Self {
79        Self {
80            semaphore: Arc::new(Semaphore::new(limit)),
81            state: Mutex::new(WorkerLimitState {
82                limit,
83                capacity: limit,
84            }),
85            active: AtomicU32::new(0),
86        }
87    }
88
89    fn state(&self) -> MutexGuard<'_, WorkerLimitState> {
90        self.state.lock().unwrap_or_else(PoisonError::into_inner)
91    }
92
93    async fn acquire(self: &Arc<Self>) -> Option<WorkerPermit> {
94        let permit = Arc::clone(&self.semaphore).acquire_owned().await.ok()?;
95
96        self.active.fetch_add(1, SeqCst);
97
98        Some(WorkerPermit {
99            limiter: Arc::clone(self),
100            permit: Some(permit),
101        })
102    }
103
104    fn set_limit(&self, limit: usize) {
105        let mut state = self.state();
106
107        state.limit = limit;
108
109        if limit > state.capacity {
110            self.semaphore.add_permits(limit - state.capacity);
111            state.capacity = limit;
112        } else if limit < state.capacity {
113            let forgotten = self.semaphore.forget_permits(state.capacity - limit);
114
115            state.capacity -= forgotten;
116        }
117    }
118
119    fn limit(&self) -> usize {
120        self.state().limit
121    }
122}
123
124struct WorkerPermit {
125    limiter: Arc<WorkerLimiter>,
126    permit: Option<OwnedSemaphorePermit>,
127}
128
129impl Drop for WorkerPermit {
130    fn drop(&mut self) {
131        self.limiter.active.fetch_sub(1, SeqCst);
132
133        let Some(permit) = self.permit.take() else {
134            return;
135        };
136        let mut state = self.limiter.state();
137
138        if state.capacity > state.limit {
139            permit.forget();
140            state.capacity -= 1;
141        }
142    }
143}
144
145impl WorkerPool {
146    #[must_use]
147    pub fn new(max_workers: usize, catalog: &'static ContentCatalog) -> Self {
148        Self {
149            catalog,
150            limiter: Arc::new(WorkerLimiter::new(max_workers)),
151            completed_chunks: Arc::new(AtomicU64::new(0)),
152            sims_completed: Arc::new(AtomicU64::new(0)),
153            work_tx: None,
154            result_rx: None,
155            resolver: None,
156        }
157    }
158
159    pub fn set_resolver<R>(&mut self, resolver: R)
160    where
161        R: DataResolver + Send + Sync + 'static,
162    {
163        self.resolver = Some(DynDataResolver::new_arc(resolver));
164    }
165
166    /// Starts worker tasks on `handle`.
167    ///
168    /// # Panics
169    ///
170    /// Panics if no resolver was configured before starting the pool.
171    pub fn start(&mut self, handle: &Handle) {
172        let (work_tx, mut work_rx) = mpsc::channel::<WorkBatch>(WORK_CHANNEL_SIZE);
173        let (result_tx, result_rx) = mpsc::channel::<WorkBatchResult>(RESULT_CHANNEL_SIZE);
174
175        self.work_tx = Some(work_tx);
176        self.result_rx = Some(result_rx);
177
178        let limiter = Arc::clone(&self.limiter);
179        let catalog = self.catalog;
180        let completed = Arc::clone(&self.completed_chunks);
181        let sims = Arc::clone(&self.sims_completed);
182        let resolver = self
183            .resolver
184            .clone()
185            .expect("resolver must be set before start");
186        let task_handle = handle.clone();
187
188        handle.spawn(async move {
189            while let Some(batch) = work_rx.recv().await {
190                let limiter = Arc::clone(&limiter);
191                let completed = Arc::clone(&completed);
192                let sims = Arc::clone(&sims);
193                // #t(rust_clone_in_loop) each chunk task needs its own sender handle; cloning mpsc::Sender is the standard pattern
194                let result_tx = result_tx.clone();
195                let resolver = Arc::clone(&resolver);
196                // #t(rust_clone_in_loop) each chunk task owns a Handle to spawn per-item subtasks
197                let item_handle = task_handle.clone();
198
199                task_handle.spawn(async move {
200                    let outcome =
201                        run_batch(batch, catalog, &limiter, &sims, &resolver, &item_handle).await;
202
203                    completed.fetch_add(1, SeqCst);
204                    let _ = result_tx.send(outcome).await;
205                });
206            }
207        });
208    }
209
210    pub fn result_rx(&mut self) -> Option<mpsc::Receiver<WorkBatchResult>> {
211        self.result_rx.take()
212    }
213
214    #[must_use]
215    pub fn work_tx(&self) -> Option<mpsc::Sender<WorkBatch>> {
216        self.work_tx.clone()
217    }
218
219    /// Update the maximum number of concurrently executing work items.
220    pub fn set_max_workers(&self, max_workers: usize) {
221        self.limiter.set_limit(max_workers);
222    }
223
224    #[expect(
225        clippy::cast_possible_truncation,
226        clippy::cast_precision_loss,
227        reason = "worker counters are bounded by configured core counts"
228    )]
229    #[must_use]
230    pub fn stats(&self) -> NodeStats {
231        let completed = self.completed_chunks.load(SeqCst);
232        let active = self.limiter.active.load(SeqCst);
233        let sims = self.sims_completed.load(SeqCst);
234        let max = self.limiter.limit() as u32;
235
236        NodeStats {
237            active_jobs: active,
238            completed_chunks: completed,
239            sims_per_second: if completed > 0 {
240                sims as f64 / SECONDS_PER_MINUTE
241            } else {
242                0.0
243            },
244            busy_workers: active,
245            max_workers: max,
246            total_cores: 0,
247            cpu_usage: if max > 0 {
248                // #t(rust_lossy_cast) active and max are small u32 worker counts, f32 is fine
249                active as f32 / max as f32
250            } else {
251                0.0
252            },
253        }
254    }
255
256    pub(crate) fn set_dyn_resolver(&mut self, resolver: Arc<DynDataResolver<'static>>) {
257        self.resolver = Some(resolver);
258    }
259}
260
261impl std::fmt::Debug for WorkerPool {
262    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
263        f.debug_struct("WorkerPool")
264            .field("catalog", &self.catalog)
265            .field("limiter", &self.limiter)
266            .field("completed_chunks", &self.completed_chunks)
267            .field("sims_completed", &self.sims_completed)
268            .field(
269                "resolver",
270                &self.resolver.as_ref().map(|_| "<DynDataResolver>"),
271            )
272            .finish_non_exhaustive()
273    }
274}
275
276async fn run_batch(
277    batch: WorkBatch,
278    catalog: &'static ContentCatalog,
279    limiter: &Arc<WorkerLimiter>,
280    sims: &Arc<AtomicU64>,
281    resolver: &Arc<DynDataResolver<'static>>,
282    handle: &Handle,
283) -> WorkBatchResult {
284    let start = Instant::now();
285    #[expect(
286        clippy::cast_possible_truncation,
287        reason = "chunk index intentionally uses the low 32 bits as an RNG index"
288    )]
289    let chunk_index = batch.chunk_id.get() as u32;
290    let mut handles = Vec::with_capacity(batch.items.len());
291
292    for item in batch.items {
293        // #t(block: rust_clone_in_loop) each item task owns handles; WorkContext clone is Arc-cheap
294        let limiter = Arc::clone(limiter);
295        let sims = Arc::clone(sims);
296        let resolver = Arc::clone(resolver);
297        let context = batch.context.clone();
298        let job_id = batch.job_id;
299
300        let task = handle.spawn(async move {
301            let _permit = limiter.acquire().await?;
302            let outcome = run_item(
303                catalog,
304                context.base_intent(),
305                context.payload(),
306                &item,
307                chunk_index,
308                resolver,
309            )
310            .await;
311
312            match outcome {
313                Ok(result) => {
314                    sims.fetch_add(u64::from(result.iterations), SeqCst);
315
316                    Some(result)
317                }
318                Err(error) => {
319                    log_item_failure(job_id, item.item_id, &error);
320
321                    None
322                }
323            }
324        });
325
326        handles.push(task);
327    }
328
329    let mut results = Vec::with_capacity(handles.len());
330
331    for task in handles {
332        match task.await {
333            Ok(Some(result)) => results.push(result),
334            Ok(None) => {}
335            Err(error) => log_task_failure(batch.job_id, &error),
336        }
337    }
338
339    #[expect(
340        clippy::cast_possible_truncation,
341        reason = "elapsed milliseconds fit in u64 for process-lifetime work"
342    )]
343    let elapsed_ms = start.elapsed().as_millis() as u64;
344
345    WorkBatchResult {
346        job_id: batch.job_id,
347        chunk_id: batch.chunk_id,
348        claim_token: batch.claim_token,
349        work_context_hash: batch.work_context_hash,
350        results,
351        elapsed_ms,
352    }
353}
354
355#[cfg(test)]
356mod tests {
357    use googletest::prelude::*;
358
359    use super::*;
360
361    static EMPTY_CATALOG: ContentCatalog = ContentCatalog::new(&[], &[], &[]);
362
363    #[gtest]
364    fn worker_pool_reports_updated_limit() -> Result<()> {
365        let pool = WorkerPool::new(2, &EMPTY_CATALOG);
366
367        pool.set_max_workers(5);
368
369        verify_that!(pool.stats().max_workers, eq(5))
370    }
371
372    #[gtest]
373    #[tokio::test]
374    async fn increasing_limit_unblocks_waiting_worker() -> Result<()> {
375        let limiter = Arc::new(WorkerLimiter::new(1));
376        let first = limiter.acquire().await.or_fail()?;
377        let waiting_limiter = Arc::clone(&limiter);
378        let waiting = tokio::spawn(async move { waiting_limiter.acquire().await });
379
380        tokio::task::yield_now().await;
381        verify_false!(waiting.is_finished())?;
382
383        limiter.set_limit(2);
384        let second = waiting.await.or_fail()?.or_fail()?;
385
386        verify_that!(limiter.active.load(SeqCst), eq(2))?;
387        verify_that!(limiter.limit(), eq(2))?;
388
389        drop(first);
390        drop(second);
391
392        verify_that!(limiter.active.load(SeqCst), eq(0))
393    }
394
395    #[gtest]
396    #[tokio::test]
397    async fn decreasing_limit_waits_for_excess_workers_to_finish() -> Result<()> {
398        let limiter = Arc::new(WorkerLimiter::new(2));
399        let first = limiter.acquire().await.or_fail()?;
400        let second = limiter.acquire().await.or_fail()?;
401
402        limiter.set_limit(1);
403        drop(first);
404
405        let waiting_limiter = Arc::clone(&limiter);
406        let waiting = tokio::spawn(async move { waiting_limiter.acquire().await });
407
408        tokio::task::yield_now().await;
409        verify_false!(waiting.is_finished())?;
410
411        drop(second);
412        let third = waiting.await.or_fail()?.or_fail()?;
413
414        verify_that!(limiter.active.load(SeqCst), eq(1))?;
415        verify_that!(limiter.limit(), eq(1))?;
416
417        drop(third);
418
419        verify_that!(limiter.active.load(SeqCst), eq(0))
420    }
421}