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#[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#[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
53pub 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 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 let result_tx = result_tx.clone();
195 let resolver = Arc::clone(&resolver);
196 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 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 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 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}