1use aion_core::{Event, TimerId};
4use std::collections::HashMap;
5use std::sync::Arc;
6use std::time::Duration;
7
8use aion_store::{ReadableEventStore, StoreError};
9use chrono::{DateTime, Utc};
10
11use crate::time::{TimerService, TimerServiceError};
12
13pub struct TimerRecovery {
15 store: Arc<dyn ReadableEventStore>,
16 timer_service: Arc<TimerService>,
17 recovery_interval: Duration,
18 now: fn() -> DateTime<Utc>,
19}
20
21#[derive(thiserror::Error, Debug, Clone, PartialEq, Eq)]
23pub enum TimerRecoveryError {
24 #[error("timer recovery store operation failed: {0}")]
26 Store(#[from] StoreError),
27
28 #[error("timer recovery fire operation failed: {0}")]
30 Timer(#[from] TimerServiceError),
31}
32
33impl TimerRecovery {
34 #[must_use]
36 pub fn new(
37 store: Arc<dyn ReadableEventStore>,
38 timer_service: Arc<TimerService>,
39 recovery_interval: Duration,
40 ) -> Self {
41 Self::with_clock(store, timer_service, recovery_interval, Utc::now)
42 }
43
44 #[must_use]
46 pub fn with_clock(
47 store: Arc<dyn ReadableEventStore>,
48 timer_service: Arc<TimerService>,
49 recovery_interval: Duration,
50 now: fn() -> DateTime<Utc>,
51 ) -> Self {
52 Self {
53 store,
54 timer_service,
55 recovery_interval,
56 now,
57 }
58 }
59
60 pub async fn recover_on_startup(
69 &self,
70 now: DateTime<Utc>,
71 ) -> Result<usize, TimerRecoveryError> {
72 let fired = self.recover_due(now).await?;
73 self.rearm_future_from_active_histories(now).await?;
74 Ok(fired)
75 }
76
77 pub async fn tick(&self) -> Result<usize, TimerRecoveryError> {
86 self.recover_due((self.now)()).await
87 }
88
89 #[must_use]
91 pub const fn recovery_interval(&self) -> Duration {
92 self.recovery_interval
93 }
94
95 async fn recover_due(&self, now: DateTime<Utc>) -> Result<usize, TimerRecoveryError> {
96 let due_timers = self.store.expired_timers(now).await?;
97 let count = due_timers.len();
98 for entry in due_timers {
99 self.timer_service
100 .fire_timer(entry.workflow_id, entry.timer_id, entry.fire_at)
101 .await?;
102 }
103 Ok(count)
104 }
105
106 async fn rearm_future_from_active_histories(
107 &self,
108 now: DateTime<Utc>,
109 ) -> Result<usize, TimerRecoveryError> {
110 let mut rearmed = 0;
111 for workflow_id in self.store.list_active().await? {
112 let history = self.store.read_history(&workflow_id).await?;
113 for (timer_id, fire_at) in outstanding_future_timers(&history, now) {
114 self.timer_service
115 .schedule(workflow_id.clone(), timer_id, fire_at)
116 .await?;
117 rearmed += 1;
118 }
119 }
120 Ok(rearmed)
121 }
122}
123
124fn outstanding_future_timers(
125 history: &[Event],
126 now: DateTime<Utc>,
127) -> Vec<(TimerId, DateTime<Utc>)> {
128 let mut outstanding: HashMap<TimerId, DateTime<Utc>> = HashMap::new();
129 for event in history {
130 match event {
131 Event::TimerStarted {
132 timer_id, fire_at, ..
133 } => {
134 outstanding.insert(timer_id.clone(), *fire_at);
135 }
136 Event::TimerFired { timer_id, .. } | Event::TimerCancelled { timer_id, .. } => {
137 outstanding.remove(timer_id);
138 }
139 _ => {}
140 }
141 }
142 outstanding
143 .into_iter()
144 .filter(|(_, fire_at)| *fire_at > now)
145 .collect()
146}
147
148#[cfg(test)]
149mod tests {
150 use std::sync::Arc;
151 use std::time::Duration;
152
153 use aion_core::{Event, EventEnvelope, TimerId, WorkflowId};
154 use aion_store::{InMemoryStore, ReadableEventStore, StoreError, WritableEventStore};
155 use chrono::{DateTime, Utc};
156
157 use super::{TimerRecovery, TimerRecoveryError};
158 use crate::engine_seam::test_support::{DeliveredWorkflowMessage, FakeEngineHandle};
159 use crate::engine_seam::{
160 EngineHandle, EngineSeamError, WorkflowProcessHandle, WorkflowResidency,
161 };
162 use crate::time::TimerService;
163
164 const RECOVERY_INTERVAL: Duration = Duration::from_millis(10);
165
166 #[derive(Debug, thiserror::Error)]
167 enum TestError {
168 #[error(transparent)]
169 Recovery(#[from] TimerRecoveryError),
170
171 #[error(transparent)]
172 Store(#[from] StoreError),
173
174 #[error(transparent)]
175 Engine(#[from] EngineSeamError),
176 }
177
178 fn instant(offset_seconds: i64) -> DateTime<Utc> {
179 DateTime::from_timestamp(1_700_000_000 + offset_seconds, 0).unwrap_or_default()
180 }
181
182 fn recorded_at() -> DateTime<Utc> {
183 instant(1)
184 }
185
186 fn tick_now() -> DateTime<Utc> {
187 instant(30)
188 }
189
190 fn workflow_id() -> WorkflowId {
191 WorkflowId::new_v4()
192 }
193
194 fn timer_id(sequence: u64) -> TimerId {
195 TimerId::anonymous(sequence)
196 }
197
198 fn recovery() -> (Arc<InMemoryStore>, Arc<FakeEngineHandle>, TimerRecovery) {
199 let concrete_store = Arc::new(InMemoryStore::default());
200 let writable: Arc<dyn WritableEventStore> = concrete_store.clone();
201 let readable: Arc<dyn ReadableEventStore> = concrete_store.clone();
202 let engine = Arc::new(FakeEngineHandle::recording_to(writable));
203 let timer_service = Arc::new(TimerService::with_recorded_at(
204 engine.clone(),
205 readable.clone(),
206 recorded_at,
207 ));
208 let recovery =
209 TimerRecovery::with_clock(readable, timer_service, RECOVERY_INTERVAL, tick_now);
210 (concrete_store, engine, recovery)
211 }
212
213 async fn history(
214 store: &InMemoryStore,
215 workflow_id: &WorkflowId,
216 ) -> Result<Vec<Event>, StoreError> {
217 store.read_history(workflow_id).await
218 }
219
220 fn timer_started_event(workflow_id: &WorkflowId, timer_id: &TimerId, seq: u64) -> Event {
221 Event::TimerStarted {
222 envelope: EventEnvelope {
223 seq,
224 recorded_at: instant(0),
225 workflow_id: workflow_id.clone(),
226 },
227 timer_id: timer_id.clone(),
228 fire_at: instant(5),
229 }
230 }
231
232 fn count_timer_fired(events: &[Event], timer_id: &TimerId) -> usize {
233 events
234 .iter()
235 .filter(|event| {
236 matches!(event, Event::TimerFired { timer_id: recorded, .. } if recorded == timer_id)
237 })
238 .count()
239 }
240
241 #[tokio::test]
242 async fn startup_sweep_fires_past_timer_and_delivers() -> Result<(), TestError> {
243 let process = WorkflowProcessHandle::new(42);
244 let (store, engine, recovery) = recovery();
245 let workflow_id = workflow_id();
246 let timer_id = timer_id(1);
247 let fire_at = instant(10);
248 engine.set_residency(workflow_id.clone(), WorkflowResidency::Resident(process))?;
249 engine.record_workflow_event(
250 &workflow_id,
251 timer_started_event(&workflow_id, &timer_id, 1),
252 )?;
253 store
254 .schedule_timer(&workflow_id, &timer_id, fire_at)
255 .await?;
256
257 let recovered = recovery.recover_on_startup(instant(20)).await?;
258
259 assert_eq!(recovered, 1);
260 assert_eq!(
261 count_timer_fired(&history(&store, &workflow_id).await?, &timer_id),
262 1
263 );
264 assert_eq!(
265 engine.delivered_messages()?,
266 vec![(
267 process,
268 DeliveredWorkflowMessage::TimerFired {
269 timer_id: timer_id.clone(),
270 fire_at
271 }
272 )]
273 );
274 Ok(())
275 }
276
277 #[tokio::test]
278 async fn startup_sweep_does_not_fire_future_timer() -> Result<(), TestError> {
279 let process = WorkflowProcessHandle::new(42);
280 let (store, engine, recovery) = recovery();
281 let workflow_id = workflow_id();
282 let timer_id = timer_id(2);
283 engine.set_residency(workflow_id.clone(), WorkflowResidency::Resident(process))?;
284 store
285 .schedule_timer(&workflow_id, &timer_id, instant(30))
286 .await?;
287
288 let recovered = recovery.recover_on_startup(instant(20)).await?;
289
290 assert_eq!(recovered, 0);
291 assert_eq!(
292 count_timer_fired(&history(&store, &workflow_id).await?, &timer_id),
293 0
294 );
295 assert!(engine.delivered_messages()?.is_empty());
296 Ok(())
297 }
298
299 #[tokio::test]
300 async fn tick_uses_injected_clock_and_fires_newly_due_timer_once() -> Result<(), TestError> {
301 let process = WorkflowProcessHandle::new(42);
302 let (store, engine, recovery) = recovery();
303 let workflow_id = workflow_id();
304 let timer_id = timer_id(3);
305 let fire_at = instant(25);
306 engine.set_residency(workflow_id.clone(), WorkflowResidency::Resident(process))?;
307 engine.record_workflow_event(
308 &workflow_id,
309 timer_started_event(&workflow_id, &timer_id, 1),
310 )?;
311 store
312 .schedule_timer(&workflow_id, &timer_id, fire_at)
313 .await?;
314
315 assert_eq!(recovery.recovery_interval(), RECOVERY_INTERVAL);
316 assert_eq!(recovery.tick().await?, 1);
317 assert_eq!(recovery.tick().await?, 1);
318
319 assert_eq!(
320 count_timer_fired(&history(&store, &workflow_id).await?, &timer_id),
321 1
322 );
323 assert_eq!(engine.delivered_messages()?.len(), 1);
324 Ok(())
325 }
326
327 #[tokio::test]
328 async fn running_startup_sweep_twice_fires_due_timer_once_total() -> Result<(), TestError> {
329 let process = WorkflowProcessHandle::new(42);
330 let (store, engine, recovery) = recovery();
331 let workflow_id = workflow_id();
332 let timer_id = timer_id(4);
333 let fire_at = instant(10);
334 engine.set_residency(workflow_id.clone(), WorkflowResidency::Resident(process))?;
335 engine.record_workflow_event(
336 &workflow_id,
337 timer_started_event(&workflow_id, &timer_id, 1),
338 )?;
339 store
340 .schedule_timer(&workflow_id, &timer_id, fire_at)
341 .await?;
342
343 recovery.recover_on_startup(instant(20)).await?;
344 recovery.recover_on_startup(instant(20)).await?;
345
346 assert_eq!(
347 count_timer_fired(&history(&store, &workflow_id).await?, &timer_id),
348 1
349 );
350 assert_eq!(engine.delivered_messages()?.len(), 1);
351 Ok(())
352 }
353
354 #[tokio::test]
355 async fn cancelled_timer_is_never_fired_by_recovery() -> Result<(), TestError> {
356 let process = WorkflowProcessHandle::new(42);
357 let (store, engine, recovery) = recovery();
358 let workflow_id = workflow_id();
359 let timer_id = timer_id(5);
360 let fire_at = instant(10);
361 engine.set_residency(workflow_id.clone(), WorkflowResidency::Resident(process))?;
362 store
363 .schedule_timer(&workflow_id, &timer_id, fire_at)
364 .await?;
365 engine.record_workflow_event(
366 &workflow_id,
367 Event::TimerCancelled {
368 envelope: EventEnvelope {
369 seq: 1,
370 recorded_at: instant(9),
371 workflow_id: workflow_id.clone(),
372 },
373 timer_id: timer_id.clone(),
374 },
375 )?;
376
377 recovery.recover_on_startup(instant(20)).await?;
378
379 assert_eq!(
380 count_timer_fired(&history(&store, &workflow_id).await?, &timer_id),
381 0
382 );
383 assert!(engine.delivered_messages()?.is_empty());
384 Ok(())
385 }
386}