Skip to main content

pgtask_worker/
registry.rs

1use std::{collections::HashMap, future::Future, pin::Pin, sync::Arc};
2
3use pgtask_core::{
4    EnqueueRequest, HandlerVersion, LeaseToken, RetryPolicy, SignalName, StepName, Task, TaskId, TaskName,
5};
6use pgtask_postgres::{ResultWait, ResultWaitRequest, SignalWait, SignalWaitRequest, SpawnRequest, Store};
7use serde_json::{Value, json};
8use thiserror::Error;
9use tokio_util::sync::CancellationToken;
10use tracing::{Instrument, info_span};
11
12pub type HandlerFuture = Pin<Box<dyn Future<Output = Result<Value, HandlerError>> + Send>>;
13
14type HandlerFunction = dyn Fn(Task, TaskContext) -> HandlerFuture + Send + Sync;
15
16#[derive(Clone, Debug, Error)]
17#[error("task handler failed")]
18pub struct HandlerError {
19    pub error: Value,
20    pub retryable: bool,
21    control: HandlerControl,
22}
23
24#[derive(Clone, Copy, Debug, Eq, PartialEq)]
25enum HandlerControl {
26    Failure,
27    Suspended,
28}
29
30impl HandlerError {
31    pub fn retryable(message: impl Into<String>) -> Self {
32        Self {
33            error: json!({"type": "handler_error", "message": message.into()}),
34            retryable: true,
35            control: HandlerControl::Failure,
36        }
37    }
38
39    pub fn terminal(message: impl Into<String>) -> Self {
40        Self {
41            error: json!({"type": "handler_error", "message": message.into()}),
42            retryable: false,
43            control: HandlerControl::Failure,
44        }
45    }
46
47    fn checkpoint(kind: &'static str, message: impl Into<String>) -> Self {
48        Self {
49            error: json!({"type": kind, "message": message.into()}),
50            retryable: true,
51            control: HandlerControl::Failure,
52        }
53    }
54
55    pub fn suspended() -> Self {
56        Self {
57            error: json!({"type": "suspended"}),
58            retryable: false,
59            control: HandlerControl::Suspended,
60        }
61    }
62
63    pub fn is_suspended(&self) -> bool {
64        self.control == HandlerControl::Suspended
65    }
66}
67
68#[derive(Clone)]
69pub struct TaskContext {
70    store: Store,
71    task_id: TaskId,
72    handler_version: HandlerVersion,
73    attempt: u16,
74    lease_token: LeaseToken,
75    cancellation: CancellationToken,
76}
77
78impl TaskContext {
79    pub fn cancellation_token(&self) -> CancellationToken {
80        self.cancellation.clone()
81    }
82
83    pub async fn step<F, Fut>(&self, step_name: &StepName, occurrence: u32, operation: F) -> Result<Value, HandlerError>
84    where
85        F: FnOnce() -> Fut,
86        Fut: Future<Output = Result<Value, HandlerError>>,
87    {
88        let span = info_span!(
89            "pgtask.checkpoint",
90            pgtask.task.id = %self.task_id,
91            pgtask.step.name = %step_name,
92            pgtask.step.occurrence = occurrence,
93        );
94        async {
95            if self.cancellation.is_cancelled() {
96                return Err(HandlerError::checkpoint("lease_lost", "task lease is no longer active"));
97            }
98            if let Some(checkpoint) = self
99                .store
100                .get_checkpoint(self.task_id, self.handler_version, step_name, occurrence)
101                .await
102                .map_err(|error| HandlerError::checkpoint("checkpoint_read_error", error.to_string()))?
103            {
104                return Ok(checkpoint.value);
105            }
106            let value = operation().await?;
107            self.store
108                .commit_checkpoint(
109                    self.task_id,
110                    self.attempt,
111                    self.lease_token,
112                    step_name,
113                    occurrence,
114                    &value,
115                )
116                .await
117                .map_err(|error| HandlerError::checkpoint("checkpoint_write_error", error.to_string()))?
118                .map(|checkpoint| checkpoint.value)
119                .ok_or_else(|| HandlerError::checkpoint("lease_lost", "task lease is no longer active"))
120        }
121        .instrument(span)
122        .await
123    }
124
125    pub async fn sleep_until(
126        &self,
127        step_name: &StepName,
128        occurrence: u32,
129        wake_at: chrono::DateTime<chrono::Utc>,
130    ) -> Result<(), HandlerError> {
131        if self.checkpoint_exists(step_name, occurrence).await? {
132            return Ok(());
133        }
134        self.store
135            .sleep_until(
136                self.task_id,
137                self.attempt,
138                self.lease_token,
139                step_name,
140                occurrence,
141                wake_at,
142            )
143            .await
144            .map_err(|error| HandlerError::checkpoint("sleep_write_error", error.to_string()))?
145            .ok_or_else(|| HandlerError::checkpoint("lease_lost", "task lease is no longer active"))?;
146        Err(HandlerError::suspended())
147    }
148
149    pub async fn sleep_for(
150        &self,
151        step_name: &StepName,
152        occurrence: u32,
153        duration: std::time::Duration,
154    ) -> Result<(), HandlerError> {
155        if self.checkpoint_exists(step_name, occurrence).await? {
156            return Ok(());
157        }
158        self.store
159            .sleep_for(
160                self.task_id,
161                self.attempt,
162                self.lease_token,
163                step_name,
164                occurrence,
165                duration,
166            )
167            .await
168            .map_err(|error| HandlerError::checkpoint("sleep_write_error", error.to_string()))?
169            .ok_or_else(|| HandlerError::checkpoint("lease_lost", "task lease is no longer active"))?;
170        Err(HandlerError::suspended())
171    }
172
173    pub async fn wait_for_signal(
174        &self,
175        step_name: &StepName,
176        occurrence: u32,
177        signal_name: &SignalName,
178        signal_occurrence: u32,
179        timeout: Option<std::time::Duration>,
180    ) -> Result<Option<Value>, HandlerError> {
181        if let Some(checkpoint) = self
182            .store
183            .get_checkpoint(self.task_id, self.handler_version, step_name, occurrence)
184            .await
185            .map_err(|error| HandlerError::checkpoint("checkpoint_read_error", error.to_string()))?
186        {
187            return decode_signal_checkpoint(&checkpoint.value);
188        }
189        match self
190            .store
191            .wait_for_signal(SignalWaitRequest {
192                task_id: self.task_id,
193                attempt: self.attempt,
194                lease_token: self.lease_token,
195                step_name,
196                occurrence,
197                signal_name,
198                signal_occurrence,
199                timeout,
200            })
201            .await
202            .map_err(|error| HandlerError::checkpoint("signal_wait_error", error.to_string()))?
203        {
204            Some(SignalWait::Ready(checkpoint)) => decode_signal_checkpoint(&checkpoint),
205            Some(SignalWait::Waiting) => Err(HandlerError::suspended()),
206            None => Err(HandlerError::checkpoint("lease_lost", "task lease is no longer active")),
207        }
208    }
209
210    pub async fn wait_for_result(
211        &self,
212        step_name: &StepName,
213        occurrence: u32,
214        result_task_id: TaskId,
215        timeout: Option<std::time::Duration>,
216    ) -> Result<Value, HandlerError> {
217        if let Some(checkpoint) = self
218            .store
219            .get_checkpoint(self.task_id, self.handler_version, step_name, occurrence)
220            .await
221            .map_err(|error| HandlerError::checkpoint("checkpoint_read_error", error.to_string()))?
222        {
223            return Ok(checkpoint.value);
224        }
225        match self
226            .store
227            .wait_for_result(ResultWaitRequest {
228                task_id: self.task_id,
229                attempt: self.attempt,
230                lease_token: self.lease_token,
231                step_name,
232                occurrence,
233                result_task_id,
234                timeout,
235            })
236            .await
237            .map_err(|error| HandlerError::checkpoint("result_wait_error", error.to_string()))?
238        {
239            Some(ResultWait::Ready(checkpoint)) => Ok(checkpoint),
240            Some(ResultWait::Waiting) => Err(HandlerError::suspended()),
241            None => Err(HandlerError::checkpoint("lease_lost", "task lease is no longer active")),
242        }
243    }
244
245    pub async fn spawn(
246        &self,
247        step_name: &StepName,
248        occurrence: u32,
249        request: &EnqueueRequest,
250    ) -> Result<TaskId, HandlerError> {
251        self.store
252            .spawn_task(SpawnRequest {
253                parent_task_id: self.task_id,
254                parent_attempt: self.attempt,
255                parent_lease_token: self.lease_token,
256                step_name,
257                occurrence,
258                task: request,
259            })
260            .await
261            .map_err(|error| HandlerError::checkpoint("child_spawn_error", error.to_string()))?
262            .map(|result| result.task_id)
263            .ok_or_else(|| HandlerError::checkpoint("lease_lost", "task lease is no longer active"))
264    }
265
266    async fn checkpoint_exists(&self, step_name: &StepName, occurrence: u32) -> Result<bool, HandlerError> {
267        self.store
268            .get_checkpoint(self.task_id, self.handler_version, step_name, occurrence)
269            .await
270            .map(|checkpoint| checkpoint.is_some())
271            .map_err(|error| HandlerError::checkpoint("checkpoint_read_error", error.to_string()))
272    }
273
274    pub(crate) fn new(store: Store, task: &Task, lease_token: LeaseToken, cancellation: CancellationToken) -> Self {
275        Self {
276            store,
277            task_id: task.id,
278            handler_version: task.handler_version,
279            attempt: task.attempt,
280            lease_token,
281            cancellation,
282        }
283    }
284}
285
286fn decode_signal_checkpoint(checkpoint: &Value) -> Result<Option<Value>, HandlerError> {
287    let Some(checkpoint) = checkpoint.as_object() else {
288        return Err(HandlerError::terminal("signal checkpoint is not an object"));
289    };
290    match checkpoint.get("outcome").and_then(Value::as_str) {
291        Some("signal") => checkpoint
292            .get("value")
293            .cloned()
294            .map(Some)
295            .ok_or_else(|| HandlerError::terminal("signal checkpoint has no value")),
296        Some("timeout") => Ok(None),
297        _ => Err(HandlerError::terminal("signal checkpoint has an invalid outcome")),
298    }
299}
300
301#[derive(Clone)]
302pub(crate) struct RegisteredHandler {
303    pub function: Arc<HandlerFunction>,
304    pub retry_policy: RetryPolicy,
305}
306
307#[derive(Clone, Default)]
308pub struct HandlerRegistry {
309    handlers: HashMap<(TaskName, HandlerVersion), RegisteredHandler>,
310}
311
312impl HandlerRegistry {
313    pub fn new() -> Self {
314        Self::default()
315    }
316
317    pub fn register<F, Fut>(
318        &mut self,
319        task_name: TaskName,
320        handler_version: HandlerVersion,
321        retry_policy: RetryPolicy,
322        handler: F,
323    ) -> bool
324    where
325        F: Fn(Task) -> Fut + Send + Sync + 'static,
326        Fut: Future<Output = Result<Value, HandlerError>> + Send + 'static,
327    {
328        let registered = RegisteredHandler {
329            function: Arc::new(move |task, _context| Box::pin(handler(task))),
330            retry_policy,
331        };
332        self.handlers.insert((task_name, handler_version), registered).is_none()
333    }
334
335    pub fn register_durable<F, Fut>(
336        &mut self,
337        task_name: TaskName,
338        handler_version: HandlerVersion,
339        retry_policy: RetryPolicy,
340        handler: F,
341    ) -> bool
342    where
343        F: Fn(Task, TaskContext) -> Fut + Send + Sync + 'static,
344        Fut: Future<Output = Result<Value, HandlerError>> + Send + 'static,
345    {
346        let registered = RegisteredHandler {
347            function: Arc::new(move |task, context| Box::pin(handler(task, context))),
348            retry_policy,
349        };
350        self.handlers.insert((task_name, handler_version), registered).is_none()
351    }
352
353    pub fn capabilities(&self) -> Vec<(TaskName, HandlerVersion)> {
354        self.handlers.keys().cloned().collect()
355    }
356
357    pub(crate) fn registrations(&self) -> Vec<(TaskName, HandlerVersion, RetryPolicy)> {
358        self.handlers
359            .iter()
360            .map(|((task_name, handler_version), handler)| (task_name.clone(), *handler_version, handler.retry_policy))
361            .collect()
362    }
363
364    pub(crate) fn get(&self, task_name: &TaskName, handler_version: HandlerVersion) -> Option<&RegisteredHandler> {
365        self.handlers.get(&(task_name.clone(), handler_version))
366    }
367}