use std::{collections::HashMap, fmt, sync::Arc, time::Duration};
use futures_util::{FutureExt, future::BoxFuture};
use log::{error, info};
use serde::{Serialize, de::DeserializeOwned};
use sqlx::PgPool;
use tokio::{task::JoinSet, time::sleep};
use uuid::Uuid;
use crate::{
context::TaskContext,
db::{claim_tasks, duration_seconds},
error::{Error, Result},
execution::SharedExecutionService,
executor::{ExecutionContext, execute_task},
metrics::QueueMetrics,
queue::Queue,
task::{Task, validate_task_name},
types::Json,
};
const DEFAULT_POLL_INTERVAL: Duration = Duration::from_millis(250);
pub(crate) type ErasedTaskExecutor =
Arc<dyn Fn(Json, TaskContext) -> BoxFuture<'static, Result<Json>> + Send + Sync>;
pub(crate) struct RegisteredTask {
pub executor: ErasedTaskExecutor,
}
impl fmt::Debug for RegisteredTask {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("RegisteredTask").field("executor", &"<task executor>").finish()
}
}
#[derive(Clone, Copy, Debug)]
struct WorkerRuntime {
lease_duration: Duration,
concurrency: usize,
}
impl Default for WorkerRuntime {
fn default() -> Self {
Self { lease_duration: Duration::from_secs(120), concurrency: 1 }
}
}
async fn run_worker_loop<S>(
pool: PgPool,
queue_name: String,
registry: Arc<HashMap<String, RegisteredTask>>,
metrics: QueueMetrics,
execution_service: SharedExecutionService,
runtime: WorkerRuntime,
shutdown: impl Future<Output = S> + Send,
) -> Result<()>
where
S: Send,
{
let worker_id = default_worker_id();
let lease_seconds = worker_lease_seconds(runtime.lease_duration)?;
let supported_tasks = registered_task_names(®istry);
let log_queue_name = queue_name.as_str();
let log_worker_id = worker_id.as_str();
let mut executing: JoinSet<Result<()>> = JoinSet::new();
let mut terminal_error = None;
let execution_context = ExecutionContext::new(
pool.clone(),
queue_name.clone(),
Arc::clone(®istry),
metrics.clone(),
execution_service,
);
tokio::pin!(shutdown);
loop {
while let Some(joined) = executing.try_join_next() {
if let Some(err) = terminal_join_error(joined) {
terminal_error = Some(err);
break;
}
}
if terminal_error.is_some() {
break;
}
let available = runtime.concurrency.saturating_sub(executing.len());
if available == 0 {
tokio::select! {
_ = &mut shutdown => {
info!("Worker shutting down (queue={log_queue_name}, worker_id={log_worker_id})");
break;
}
joined = executing.join_next() => {
if let Some(joined) = joined
&& let Some(err) = terminal_join_error(joined)
{
terminal_error = Some(err);
break;
}
}
}
continue;
}
if shutdown.as_mut().now_or_never().is_some() {
info!("Worker shutting down (queue={log_queue_name}, worker_id={log_worker_id})");
break;
}
let batch_size = i32::try_from(available).map_err(|_| {
Error::InvalidOptions("worker concurrency exceeds PostgreSQL integer range".to_owned())
})?;
let tasks = match claim_tasks(
&pool,
&queue_name,
&worker_id,
lease_seconds,
batch_size,
&supported_tasks,
)
.await
{
Ok(tasks) => {
metrics.record_claimed(tasks.len());
tasks
}
Err(err) => {
metrics.record_claim_error();
if !is_transient_worker_error(&err) {
terminal_error = Some(err);
break;
}
error!(
"Transient worker claim error (queue={log_queue_name}, worker_id={log_worker_id}): {err:?}"
);
tokio::select! {
_ = &mut shutdown => {
info!("Worker shutting down (queue={log_queue_name}, worker_id={log_worker_id})");
break;
}
_ = sleep(DEFAULT_POLL_INTERVAL) => {}
}
continue;
}
};
if tasks.is_empty() {
tokio::select! {
_ = &mut shutdown => {
info!("Worker shutting down (queue={log_queue_name}, worker_id={log_worker_id})");
break;
}
_ = sleep(DEFAULT_POLL_INTERVAL) => {}
}
continue;
}
for task in tasks {
let execution_context = execution_context.clone();
let queue_name = queue_name.clone();
let worker_id = worker_id.clone();
executing.spawn(async move {
match execute_task(execution_context, task, lease_seconds).await {
Ok(())
| Err(
Error::Suspended
| Error::Cancelled
| Error::FailedRun
| Error::LeaseLost,
) => Ok(()),
Err(err) if is_transient_worker_error(&err) => {
error!(
"Transient task execution infrastructure error (queue={queue_name}, worker_id={worker_id}): {err:?}"
);
Ok(())
}
Err(err) => Err(err),
}
});
}
}
while let Some(joined) = executing.join_next().await {
if let Some(err) = terminal_join_error(joined)
&& terminal_error.is_none()
{
terminal_error = Some(err);
}
}
if let Some(err) = terminal_error {
return Err(err);
}
Ok(())
}
fn terminal_join_error(
joined: std::result::Result<Result<()>, tokio::task::JoinError>,
) -> Option<Error> {
match joined {
Ok(Ok(())) => None,
Ok(Err(err)) => Some(err),
Err(err) => Some(Error::Other(format!("worker task join failed: {err}"))),
}
}
fn is_transient_worker_error(error: &Error) -> bool {
let Error::Database(error) = error else {
return false;
};
match error {
sqlx::Error::Io(_) | sqlx::Error::Tls(_) | sqlx::Error::PoolTimedOut => true,
sqlx::Error::Database(database) => {
database.code().is_some_and(|code| is_transient_sqlstate(code.as_ref()))
}
_ => false,
}
}
fn is_transient_sqlstate(code: &str) -> bool {
code.starts_with("08")
|| code.starts_with("40")
|| code.starts_with("53")
|| matches!(code, "55P03" | "57P01" | "57P02" | "57P03")
}
pub trait TaskExecutor<Input, Output>: Send + Sync + 'static {
fn execute(
&self,
input: Input,
context: TaskContext,
) -> impl Future<Output = Result<Output>> + Send;
}
pub trait TaskHandler<Input, Output>:
Fn(Input, TaskContext) -> <Self as TaskHandler<Input, Output>>::Future + Send + Sync + 'static
{
type Future: Future<Output = Result<Output>> + Send + 'static;
}
impl<Input, Output, F, Fut> TaskHandler<Input, Output> for F
where
F: Fn(Input, TaskContext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Output>> + Send + 'static,
{
type Future = Fut;
}
impl<Input, Output, F> TaskExecutor<Input, Output> for F
where
F: TaskHandler<Input, Output>,
{
fn execute(
&self,
input: Input,
context: TaskContext,
) -> impl Future<Output = Result<Output>> + Send {
(self)(input, context)
}
}
#[derive(Debug)]
pub struct WorkerBuilder {
queue: Queue,
registry: HashMap<String, RegisteredTask>,
runtime: WorkerRuntime,
error: Option<Error>,
}
impl WorkerBuilder {
pub(crate) fn new(queue: Queue) -> Self {
Self { queue, registry: HashMap::new(), runtime: WorkerRuntime::default(), error: None }
}
#[must_use]
pub fn lease_duration(mut self, lease_duration: Duration) -> Self {
if self.error.is_some() {
return self;
}
match duration_seconds(lease_duration) {
Ok(0) => {
self.error = Some(Error::InvalidOptions(
"lease duration must round to at least 1 second".to_owned(),
));
return self;
}
Ok(_) => {}
Err(error) => {
self.error = Some(error);
return self;
}
}
self.runtime.lease_duration = lease_duration;
self
}
#[must_use]
pub fn concurrency(mut self, concurrency: usize) -> Self {
if self.error.is_some() {
return self;
}
if concurrency == 0 {
self.error =
Some(Error::InvalidOptions("worker concurrency must be at least 1".to_owned()));
return self;
}
if i32::try_from(concurrency).is_err() {
self.error = Some(Error::InvalidOptions(
"worker concurrency exceeds PostgreSQL integer range".to_owned(),
));
return self;
}
self.runtime.concurrency = concurrency;
self
}
#[must_use]
pub fn task<Input, Output>(
self,
task: Task<Input, Output>,
handler: impl TaskHandler<Input, Output>,
) -> Self
where
Input: DeserializeOwned + Send + 'static,
Output: Serialize + Send + 'static,
{
self.register_task_executor(task, handler)
}
#[must_use]
pub fn task_executor<Input, Output>(
self,
task: Task<Input, Output>,
executor: impl TaskExecutor<Input, Output>,
) -> Self
where
Input: DeserializeOwned + Send + 'static,
Output: Serialize + Send + 'static,
{
self.register_task_executor(task, executor)
}
fn register_task_executor<Input, Output>(
mut self,
task: Task<Input, Output>,
executor: impl TaskExecutor<Input, Output>,
) -> Self
where
Input: DeserializeOwned + Send + 'static,
Output: Serialize + Send + 'static,
{
if self.error.is_some() {
return self;
}
if let Err(error) = validate_task_name(task.name()) {
self.error = Some(error);
return self;
}
if self.registry.contains_key(task.name()) {
self.error = Some(Error::InvalidOptions(format!(
"task {:?} is already registered",
task.name()
)));
return self;
}
let executor = Arc::new(executor);
let erased: ErasedTaskExecutor = Arc::new(move |raw, context| {
let executor = Arc::clone(&executor);
Box::pin(async move {
let input = serde_json::from_value::<Input>(raw)?;
let output = executor.execute(input, context).await?;
Ok(serde_json::to_value(output)?)
})
});
self.registry.insert(task.name().to_owned(), RegisteredTask { executor: erased });
self
}
pub fn build(self) -> Result<Worker> {
if let Some(error) = self.error {
return Err(error);
}
if self.registry.is_empty() {
return Err(Error::InvalidOptions("worker requires at least one task".to_owned()));
}
Ok(Worker { queue: self.queue, registry: Arc::new(self.registry), runtime: self.runtime })
}
}
#[derive(Debug)]
pub struct Worker {
queue: Queue,
registry: Arc<HashMap<String, RegisteredTask>>,
runtime: WorkerRuntime,
}
impl Worker {
pub async fn run(&self) -> Result<()> {
self.run_until(std::future::pending::<()>()).await
}
pub async fn run_until<S>(&self, shutdown: impl Future<Output = S> + Send) -> Result<()>
where
S: Send,
{
run_worker_loop(
self.queue.pool().clone(),
self.queue.name().to_owned(),
Arc::clone(&self.registry),
self.queue.metrics(),
self.queue.execution().clone(),
self.runtime,
shutdown,
)
.await
}
pub fn metrics(&self) -> QueueMetrics {
self.queue.metrics()
}
}
fn registered_task_names(registry: &HashMap<String, RegisteredTask>) -> Vec<String> {
registry.keys().cloned().collect()
}
fn default_worker_id() -> String {
format!("worker:{}", Uuid::now_v7())
}
fn worker_lease_seconds(lease_duration: Duration) -> Result<i32> {
duration_seconds(lease_duration)
}
#[cfg(test)]
mod tests {
use super::is_transient_sqlstate;
#[test]
fn transient_sqlstates_are_narrowly_classified() {
for code in ["08006", "40001", "40P01", "53300", "55P03", "57P01", "57P02", "57P03"] {
assert!(is_transient_sqlstate(code), "{code} should be transient");
}
for code in ["22023", "23505", "42703", "42P01", "ST001"] {
assert!(!is_transient_sqlstate(code), "{code} should be terminal");
}
}
}