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}