Skip to main content

mesofact_core/proxy/
worker_pool.rs

1//! WorkerPool — manages N Bun render-pool workers.
2//!
3//! Responsibilities:
4//! - Spawn N workers (default = `num_cpus / 2`, min 1) on a manifest.
5//! - Watchdog task: ping every 30 s; kill + respawn on missed pong.
6//! - Rolling reload: spawn parallel new pool, drain old.
7//! - `get()`: returns any live worker (round-robin in P9+; index 0 in P7
8//!   since Mode 2 is stubbed and no concurrent renders flow through).
9
10use crate::proxy::metrics::Metrics;
11use crate::proxy::worker_client::{WorkerClient, WorkerError};
12use std::path::PathBuf;
13use std::sync::atomic::{AtomicUsize, Ordering};
14use std::sync::{Arc, Mutex};
15use std::time::Duration;
16use tokio::sync::RwLock;
17use tracing::{error, info, warn};
18
19const PING_INTERVAL: Duration = Duration::from_secs(30);
20
21pub struct WorkerPool {
22    workers: Arc<RwLock<Vec<Arc<WorkerClient>>>>,
23    worker_entry: PathBuf,
24    manifest_path: PathBuf,
25    /// `mesofact.config.toml` path passed to each worker for adapter
26    /// registration (sqlite/r2 sources). `None` = no sources declared.
27    config_path: Option<PathBuf>,
28    #[allow(dead_code)]
29    n: usize,
30    /// Round-robin cursor for `get()` so concurrent Mode 2 renders spread
31    /// across workers instead of all serializing on worker 0's socket mutex.
32    next: AtomicUsize,
33    /// Shared metrics registry (set via `attach_metrics`); the watchdog bumps
34    /// the `restarting` gauge around a respawn. `None` in tests.
35    metrics: Mutex<Option<Arc<Metrics>>>,
36    _tmp: Arc<tempfile::TempDir>,
37}
38
39impl WorkerPool {
40    /// Spawn `n` workers loading `manifest_path` and start the watchdog.
41    pub async fn spawn(
42        manifest_json: &[u8],
43        worker_entry: PathBuf,
44        n: usize,
45    ) -> Result<Arc<Self>, WorkerError> {
46        Self::spawn_with_config(manifest_json, worker_entry, n, None).await
47    }
48
49    /// Like `spawn`, but passes `config_path` to each worker so its render
50    /// entrypoints can reach adapters declared in `mesofact.config.toml`.
51    pub async fn spawn_with_config(
52        manifest_json: &[u8],
53        worker_entry: PathBuf,
54        n: usize,
55        config_path: Option<PathBuf>,
56    ) -> Result<Arc<Self>, WorkerError> {
57        let tmp = Arc::new(tempfile::tempdir().map_err(WorkerError::Io)?);
58        let manifest_path = tmp.path().join("manifest.json");
59        tokio::fs::write(&manifest_path, manifest_json)
60            .await
61            .map_err(WorkerError::Io)?;
62
63        let mut workers = Vec::with_capacity(n);
64        for i in 0..n {
65            let sock = tmp.path().join(format!("worker-{i}.sock"));
66            info!("spawning worker {i}");
67            workers.push(Arc::new(
68                WorkerClient::spawn(sock, &manifest_path, &worker_entry, config_path.as_deref())
69                    .await?,
70            ));
71        }
72
73        let pool = Arc::new(Self {
74            workers: Arc::new(RwLock::new(workers)),
75            worker_entry,
76            manifest_path,
77            config_path,
78            n,
79            next: AtomicUsize::new(0),
80            metrics: Mutex::new(None),
81            _tmp: tmp,
82        });
83
84        pool.clone().start_watchdog();
85        Ok(pool)
86    }
87
88    /// Attach the shared metrics registry so respawns bump the `restarting`
89    /// worker-pool gauge. Called once after both the pool and `AppState` exist.
90    pub fn attach_metrics(&self, metrics: Arc<Metrics>) {
91        *self.metrics.lock().unwrap() = Some(metrics);
92    }
93
94    /// Live worker count — the `ready` worker-pool gauge, read at scrape time.
95    pub async fn live_count(&self) -> usize {
96        self.workers.read().await.len()
97    }
98
99    fn metrics(&self) -> Option<Arc<Metrics>> {
100        self.metrics.lock().unwrap().clone()
101    }
102
103    /// Return a live worker, round-robin across the pool. Per-worker I/O still
104    /// serializes on that worker's socket mutex (see `worker_client` cleanup),
105    /// so spreading requests across workers is what buys real concurrency.
106    pub async fn get(&self) -> Option<Arc<WorkerClient>> {
107        let workers = self.workers.read().await;
108        if workers.is_empty() {
109            return None;
110        }
111        let i = self.next.fetch_add(1, Ordering::Relaxed) % workers.len();
112        workers.get(i).cloned()
113    }
114
115    /// Drain all workers and wait for them to exit (used during rolling reload).
116    pub async fn drain_all(self: Arc<Self>) {
117        let workers = self.workers.read().await.clone();
118        let mut joins = Vec::with_capacity(workers.len());
119        for w in workers {
120            joins.push(tokio::spawn(async move {
121                if let Err(e) = w.drain().await {
122                    warn!("drain error: {e}");
123                }
124                let _ = w.wait().await;
125            }));
126        }
127        for j in joins {
128            let _ = j.await;
129        }
130    }
131
132    fn start_watchdog(self: Arc<Self>) {
133        tokio::spawn(async move {
134            let mut ticker = tokio::time::interval(PING_INTERVAL);
135            ticker.tick().await; // skip immediate first tick
136            loop {
137                ticker.tick().await;
138                self.watchdog_cycle().await;
139            }
140        });
141    }
142
143    async fn watchdog_cycle(&self) {
144        let snapshot: Vec<Arc<WorkerClient>> = self.workers.read().await.clone();
145        for (i, w) in snapshot.iter().enumerate() {
146            match w.ping().await {
147                Ok(()) => {}
148                Err(WorkerError::PongTimeout | WorkerError::Closed | WorkerError::Io(_)) => {
149                    warn!("worker {i} missed pong — respawning");
150                    let _ = w.kill().await;
151                    let metrics = self.metrics();
152                    if let Some(m) = &metrics {
153                        m.restarting_inc();
154                    }
155                    match self.respawn(i).await {
156                        Ok(new_w) => {
157                            let mut lock = self.workers.write().await;
158                            if i < lock.len() {
159                                lock[i] = Arc::new(new_w);
160                            }
161                            info!("worker {i} respawned");
162                        }
163                        Err(e) => error!("failed to respawn worker {i}: {e}"),
164                    }
165                    if let Some(m) = &metrics {
166                        m.restarting_dec();
167                    }
168                }
169                Err(e) => warn!("worker {i} ping error: {e}"),
170            }
171        }
172    }
173
174    async fn respawn(&self, idx: usize) -> Result<WorkerClient, WorkerError> {
175        let sock = self
176            ._tmp
177            .path()
178            .join(format!("worker-{idx}-r{}.sock", unix_now()));
179        WorkerClient::spawn(sock, &self.manifest_path, &self.worker_entry, self.config_path.as_deref())
180            .await
181    }
182}
183
184fn unix_now() -> u64 {
185    std::time::SystemTime::now()
186        .duration_since(std::time::UNIX_EPOCH)
187        .map(|d| d.as_secs())
188        .unwrap_or(0)
189}
190
191/// Spawn a new pool from fresh manifest bytes and then drain the old pool.
192/// The caller swaps the `Arc<WorkerPool>` reference atomically before draining.
193pub async fn rolling_reload(
194    old: Arc<WorkerPool>,
195    manifest_json: &[u8],
196    worker_entry: PathBuf,
197    n: usize,
198) -> Result<Arc<WorkerPool>, WorkerError> {
199    let config_path = old.config_path.clone();
200    let new_pool = WorkerPool::spawn_with_config(manifest_json, worker_entry, n, config_path).await?;
201    tokio::spawn(old.drain_all());
202    Ok(new_pool)
203}