Skip to main content

wowlab_sentinel/scheduler/
driver.rs

1//! Cancellable `PostgreSQL` notification and polling driver.
2
3use 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}