wowlab_sentinel/scheduler/
driver.rs1use std::{sync::Arc, time::Duration};
4
5use async_trait::async_trait;
6use sqlx::postgres::PgListener;
7use tokio_util::sync::CancellationToken;
8
9use super::assignment::assignment_workflow;
10use crate::state::ServerState;
11
12const NOTIFY_DEBOUNCE_MS: u64 = 50;
13
14#[derive(Clone, Debug, Eq, PartialEq)]
15enum SchedulerTrigger {
16 Notification(String),
17 PollTimeout,
18}
19
20#[async_trait]
21trait SchedulerConnection: Send {
22 async fn next(&mut self, poll_timeout: Duration) -> Result<SchedulerTrigger, sqlx::Error>;
23}
24
25#[async_trait]
26trait SchedulerConnector: Send + Sync {
27 type Connection: SchedulerConnection;
28
29 async fn connect(&self) -> Result<Self::Connection, sqlx::Error>;
30}
31
32#[async_trait]
33trait AssignmentRounds: Send + Sync {
34 fn touch(&self);
35 async fn process(&self);
36}
37
38#[async_trait]
39trait ReconnectSleeper: Send + Sync {
40 async fn sleep(&self, duration: Duration);
41}
42
43struct DriverConfig {
44 poll_timeout: Duration,
45 reconnect_delay: Duration,
46}
47
48struct DriverDependencies<C, R, S> {
49 connector: C,
50 rounds: R,
51 sleeper: S,
52}
53
54struct PgConnector<'a> {
55 pool: &'a sqlx::PgPool,
56}
57
58struct PgConnection {
59 listener: PgListener,
60}
61
62impl PgConnector<'_> {
63 async fn connect_listener(&self) -> Result<PgConnection, sqlx::Error> {
64 let mut listener = PgListener::connect_with(self.pool).await?;
65
66 listener.listen("pending_job").await?;
67 tracing::info!("Listening for pending_job notifications");
68
69 Ok(PgConnection { listener })
70 }
71}
72
73#[async_trait]
74impl SchedulerConnector for PgConnector<'_> {
75 type Connection = PgConnection;
76
77 async fn connect(&self) -> Result<Self::Connection, sqlx::Error> {
78 self.connect_listener().await
79 }
80}
81
82#[async_trait]
83impl SchedulerConnection for PgConnection {
84 async fn next(&mut self, poll_timeout: Duration) -> Result<SchedulerTrigger, sqlx::Error> {
85 match tokio::time::timeout(poll_timeout, self.listener.recv()).await {
86 Ok(Ok(notification)) => {
87 let payload = notification.payload().to_owned();
88
89 tracing::info!(payload, "NOTIFY pending_job received");
90 tokio::time::sleep(Duration::from_millis(NOTIFY_DEBOUNCE_MS)).await;
91
92 Ok(SchedulerTrigger::Notification(payload))
93 }
94 Ok(Err(error)) => {
95 tracing::error!(%error, "PgListener error, reconnecting");
96
97 Err(error)
98 }
99 Err(_) => Ok(SchedulerTrigger::PollTimeout),
100 }
101 }
102}
103
104struct StateAssignmentRounds<'a> {
105 state: &'a ServerState,
106}
107
108#[async_trait]
109impl AssignmentRounds for StateAssignmentRounds<'_> {
110 fn touch(&self) {
111 self.state.touch_scheduler();
112 }
113
114 async fn process(&self) {
115 if let Err(error) = assignment_workflow(self.state).run().await {
116 if error.is_pending_fetch() {
117 tracing::error!(%error, "Failed to fetch pending jobs");
118 } else {
119 tracing::error!(%error, "Assignment failed");
120 }
121 }
122 }
123}
124
125struct TokioSleeper;
126
127#[async_trait]
128impl ReconnectSleeper for TokioSleeper {
129 async fn sleep(&self, duration: Duration) {
130 tokio::time::sleep(duration).await;
131 }
132}
133
134pub(super) async fn run(state: Arc<ServerState>, shutdown: CancellationToken) {
135 tracing::info!("Scheduler starting");
136 let config = DriverConfig {
137 poll_timeout: Duration::from_secs(state.config.scheduler_poll_timeout_secs),
138 reconnect_delay: Duration::from_secs(state.config.scheduler_reconnect_delay_secs),
139 };
140 let dependencies = DriverDependencies {
141 connector: PgConnector {
142 pool: state.dbs.get::<super::SchedulerDb>(),
143 },
144 rounds: StateAssignmentRounds { state: &state },
145 sleeper: TokioSleeper,
146 };
147
148 drive(config, dependencies, shutdown).await;
149}
150
151async fn drive<C, R, S>(
152 config: DriverConfig,
153 dependencies: DriverDependencies<C, R, S>,
154 shutdown: CancellationToken,
155) where
156 C: SchedulerConnector,
157 R: AssignmentRounds,
158 S: ReconnectSleeper,
159{
160 loop {
161 if shutdown.is_cancelled() {
162 return;
163 }
164
165 match dependencies.connector.connect().await {
166 Ok(mut connection) => {
167 dependencies.rounds.process().await;
168
169 if drive_connection(
170 &mut connection,
171 &dependencies.rounds,
172 &shutdown,
173 config.poll_timeout,
174 )
175 .await
176 {
177 return;
178 }
179 }
180 Err(error) => log_scheduler_failure(&error),
181 }
182
183 tokio::select! {
184 () = shutdown.cancelled() => return,
185 () = dependencies.sleeper.sleep(config.reconnect_delay) => {}
186 }
187 }
188}
189
190async fn drive_connection<C, R>(
191 connection: &mut C,
192 rounds: &R,
193 shutdown: &CancellationToken,
194 poll_timeout: Duration,
195) -> bool
196where
197 C: SchedulerConnection,
198 R: AssignmentRounds,
199{
200 loop {
201 rounds.touch();
202 let trigger = tokio::select! {
203 () = shutdown.cancelled() => return true,
204 result = connection.next(poll_timeout) => match result {
205 Ok(trigger) => trigger,
206 Err(error) => {
207 log_scheduler_failure(&error);
208 return false;
209 }
210 }
211 };
212
213 if trigger == SchedulerTrigger::PollTimeout {
214 log_poll_timeout(poll_timeout);
215 }
216
217 rounds.process().await;
218 }
219}
220
221fn log_scheduler_failure(error: &sqlx::Error) {
222 tracing::error!(%error, "Scheduler failed, reconnecting");
223}
224
225fn log_poll_timeout(poll_timeout: Duration) {
226 tracing::info!(
227 timeout_secs = poll_timeout.as_secs(),
228 "Scheduler poll timeout, checking for work"
229 );
230}
231
232#[cfg(test)]
233mod tests {
234 use std::{
235 collections::VecDeque,
236 sync::{
237 Mutex,
238 atomic::{AtomicUsize, Ordering},
239 },
240 };
241
242 use googletest::prelude::*;
243
244 use super::*;
245
246 struct FakeConnection {
247 events: VecDeque<Result<SchedulerTrigger, sqlx::Error>>,
248 }
249
250 #[async_trait]
251 impl SchedulerConnection for FakeConnection {
252 async fn next(&mut self, _poll_timeout: Duration) -> Result<SchedulerTrigger, sqlx::Error> {
253 match self.events.pop_front() {
254 Some(event) => event,
255 None => std::future::pending().await,
256 }
257 }
258 }
259
260 struct FakeConnector {
261 attempts: AtomicUsize,
262 connections: Mutex<VecDeque<Result<FakeConnection, sqlx::Error>>>,
263 }
264
265 #[async_trait]
266 impl SchedulerConnector for &FakeConnector {
267 type Connection = FakeConnection;
268
269 async fn connect(&self) -> Result<Self::Connection, sqlx::Error> {
270 self.attempts.fetch_add(1, Ordering::Relaxed);
271
272 self.connections
273 .lock()
274 .expect("connections lock")
275 .pop_front()
276 .unwrap_or_else(|| Err(sqlx::Error::Protocol("no connection".into())))
277 }
278 }
279
280 #[derive(Default)]
281 struct FakeRounds {
282 touches: AtomicUsize,
283 rounds: AtomicUsize,
284 }
285
286 #[async_trait]
287 impl AssignmentRounds for &FakeRounds {
288 fn touch(&self) {
289 self.touches.fetch_add(1, Ordering::Relaxed);
290 }
291
292 async fn process(&self) {
293 self.rounds.fetch_add(1, Ordering::Relaxed);
294 }
295 }
296
297 struct FakeSleeper {
298 sleeps: AtomicUsize,
299 shutdown: CancellationToken,
300 cancel_after: usize,
301 }
302
303 #[async_trait]
304 impl ReconnectSleeper for &FakeSleeper {
305 async fn sleep(&self, _duration: Duration) {
306 let sleeps = self.sleeps.fetch_add(1, Ordering::Relaxed) + 1;
307
308 if sleeps >= self.cancel_after {
309 self.shutdown.cancel();
310 }
311 }
312 }
313
314 fn config() -> DriverConfig {
315 DriverConfig {
316 poll_timeout: Duration::from_secs(30),
317 reconnect_delay: Duration::from_secs(5),
318 }
319 }
320
321 #[gtest]
322 #[tokio::test]
323 async fn cancellation_before_connect_has_no_external_side_effect() -> Result<()> {
324 let shutdown = CancellationToken::new();
325
326 shutdown.cancel();
327 let connector = FakeConnector {
328 attempts: AtomicUsize::new(0),
329 connections: Mutex::new(VecDeque::new()),
330 };
331 let rounds = FakeRounds::default();
332 let sleeper = FakeSleeper {
333 sleeps: AtomicUsize::new(0),
334 shutdown: shutdown.clone(),
335 cancel_after: 1,
336 };
337
338 drive(
339 config(),
340 DriverDependencies {
341 connector: &connector,
342 rounds: &rounds,
343 sleeper: &sleeper,
344 },
345 shutdown,
346 )
347 .await;
348
349 verify_eq!(connector.attempts.load(Ordering::Relaxed), 0)?;
350 verify_eq!(rounds.rounds.load(Ordering::Relaxed), 0)?;
351 verify_eq!(sleeper.sleeps.load(Ordering::Relaxed), 0)?;
352
353 Ok(())
354 }
355
356 #[gtest]
357 #[tokio::test]
358 async fn connection_failure_retries_then_processes_initial_notification_and_timeout_rounds()
359 -> Result<()> {
360 let shutdown = CancellationToken::new();
361 let connector = FakeConnector {
362 attempts: AtomicUsize::new(0),
363 connections: Mutex::new(VecDeque::from([
364 Err(sqlx::Error::Protocol("connect failed".into())),
365 Ok(FakeConnection {
366 events: VecDeque::from([
367 Ok(SchedulerTrigger::Notification("job".into())),
368 Ok(SchedulerTrigger::PollTimeout),
369 Err(sqlx::Error::Protocol("listener failed".into())),
370 ]),
371 }),
372 ])),
373 };
374 let rounds = FakeRounds::default();
375 let sleeper = FakeSleeper {
376 sleeps: AtomicUsize::new(0),
377 shutdown: shutdown.clone(),
378 cancel_after: 2,
379 };
380
381 drive(
382 config(),
383 DriverDependencies {
384 connector: &connector,
385 rounds: &rounds,
386 sleeper: &sleeper,
387 },
388 shutdown,
389 )
390 .await;
391
392 verify_eq!(connector.attempts.load(Ordering::Relaxed), 2)?;
393 verify_eq!(sleeper.sleeps.load(Ordering::Relaxed), 2)?;
394 verify_eq!(rounds.rounds.load(Ordering::Relaxed), 3)?;
395 verify_eq!(rounds.touches.load(Ordering::Relaxed), 3)?;
396
397 Ok(())
398 }
399}