use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use chrono::Utc;
use tracing::{debug, error, info, warn};
pub(crate) fn is_transient_db_error(msg: &str) -> bool {
let m = msg.to_ascii_lowercase();
m.contains("database is locked")
|| m.contains("database table is locked")
|| m.contains("deadlock detected")
|| m.contains("could not serialize access")
}
pub(crate) async fn retry_transient<T, E, F, Fut>(
attempts: u32,
base_delay: Duration,
mut op: F,
) -> Result<T, E>
where
E: std::fmt::Display,
F: FnMut() -> Fut,
Fut: std::future::Future<Output = Result<T, E>>,
{
let mut last: Option<E> = None;
for attempt in 1..=attempts {
match op().await {
Ok(v) => return Ok(v),
Err(e) if is_transient_db_error(&e.to_string()) && attempt < attempts => {
warn!(
attempt,
attempts,
error = %e,
"transient DB contention on task-state write — retrying"
);
tokio::time::sleep(base_delay * attempt).await;
last = Some(e);
}
Err(e) => return Err(e),
}
}
Err(last.expect("retry_transient: attempts >= 1"))
}
use crate::context::Context;
use crate::dal::DAL;
use crate::database::universal_types::UniversalUuid;
use crate::dispatcher::{ExecutionResult, TaskReadyEvent};
use crate::error::ExecutorError;
use crate::executor::types::ClaimedTask;
use crate::retry::{RetryCondition, RetryPolicy};
#[derive(Clone)]
pub struct TaskResultHandler {
dal: DAL,
total_executed: Arc<AtomicU64>,
total_failed: Arc<AtomicU64>,
runner_id: Option<UniversalUuid>,
}
impl TaskResultHandler {
pub fn new(
dal: DAL,
total_executed: Arc<AtomicU64>,
total_failed: Arc<AtomicU64>,
runner_id: Option<UniversalUuid>,
) -> Self {
Self {
dal,
total_executed,
total_failed,
runner_id,
}
}
pub async fn handle_outcome(
&self,
event: &TaskReadyEvent,
claimed_task: &ClaimedTask,
outcome: Result<Context<serde_json::Value>, ExecutorError>,
retry_policy: &RetryPolicy,
duration: Duration,
) -> ExecutionResult {
match outcome {
Ok(result_context) => {
match retry_transient(5, Duration::from_millis(100), || {
self.complete_task_transaction(claimed_task, result_context.clone_data())
})
.await
{
Ok(()) => {
self.total_executed.fetch_add(1, Ordering::SeqCst);
info!(
task_id = %event.task_execution_id,
task_name = %event.task_name,
duration_ms = duration.as_millis(),
"Task executed successfully via dispatcher"
);
ExecutionResult::success(event.task_execution_id, duration)
}
Err(e) => {
self.total_failed.fetch_add(1, Ordering::SeqCst);
let error_msg = format!("Failed to save context: {}", e);
if let Err(mark_err) =
retry_transient(5, Duration::from_millis(100), || async {
self.dal
.task_execution()
.mark_failed(
event.task_execution_id,
&error_msg,
self.runner_id,
)
.await
})
.await
{
error!(
task_id = %event.task_execution_id,
error = %mark_err,
"mark_failed FAILED after retries — task row stays Running until \
the stale-claim sweeper recovers it; workflow completion is delayed"
);
}
ExecutionResult::failure(event.task_execution_id, error_msg, duration)
}
}
}
Err(error) => {
let should_retry = self
.should_retry_task(claimed_task, &error, retry_policy)
.await
.unwrap_or(false);
if should_retry {
match retry_transient(5, Duration::from_millis(100), || {
self.schedule_task_retry(claimed_task, retry_policy)
})
.await
{
Ok(()) => {
self.total_executed.fetch_add(1, Ordering::SeqCst);
ExecutionResult::retry(
event.task_execution_id,
error.to_string(),
duration,
)
}
Err(sched_err) => {
warn!(
task_id = %event.task_execution_id,
error = %sched_err,
"Failed to schedule retry after retries — failing task \
instead of leaving it Running"
);
self.total_failed.fetch_add(1, Ordering::SeqCst);
let error_str =
format!("{} (retry scheduling failed: {})", error, sched_err);
if let Err(mark_err) =
retry_transient(5, Duration::from_millis(100), || async {
self.dal
.task_execution()
.mark_failed(
event.task_execution_id,
&error_str,
self.runner_id,
)
.await
})
.await
{
error!(
task_id = %event.task_execution_id,
error = %mark_err,
"mark_failed FAILED after retries — task row stays Running \
until the stale-claim sweeper recovers it; workflow \
completion is delayed"
);
}
ExecutionResult::failure(event.task_execution_id, error_str, duration)
}
}
} else {
self.total_failed.fetch_add(1, Ordering::SeqCst);
let error_str = error.to_string();
if let Err(mark_err) =
retry_transient(5, Duration::from_millis(100), || async {
self.dal
.task_execution()
.mark_failed(event.task_execution_id, &error_str, self.runner_id)
.await
})
.await
{
error!(
task_id = %event.task_execution_id,
error = %mark_err,
"mark_failed FAILED after retries — task row stays Running until \
the stale-claim sweeper recovers it; workflow completion is delayed"
);
}
ExecutionResult::failure(event.task_execution_id, error_str, duration)
}
}
}
}
async fn complete_task_transaction(
&self,
claimed_task: &ClaimedTask,
context: Context<serde_json::Value>,
) -> Result<(), ExecutorError> {
self.save_task_context(claimed_task, context).await?;
let applied = self
.dal
.task_execution()
.mark_completed(claimed_task.task_execution_id, self.runner_id)
.await?;
if !applied {
warn!(
task_id = %claimed_task.task_execution_id,
task_name = %claimed_task.task_name,
"Claim lost between context save and mark_completed — context row is orphaned (harmless), another runner now owns this task"
);
return Ok(());
}
info!(
task_id = %claimed_task.task_execution_id,
task_name = %claimed_task.task_name,
workflow_id = %claimed_task.workflow_execution_id,
"Task state change: -> Completed"
);
Ok(())
}
async fn save_task_context(
&self,
claimed_task: &ClaimedTask,
context: Context<serde_json::Value>,
) -> Result<(), ExecutorError> {
use crate::models::task_execution_metadata::NewTaskExecutionMetadata;
let context_id = self.dal.context().create(&context).await?;
let task_metadata_record = NewTaskExecutionMetadata {
task_execution_id: claimed_task.task_execution_id,
workflow_execution_id: claimed_task.workflow_execution_id,
task_name: claimed_task.task_name.clone(),
context_id,
};
self.dal
.task_execution_metadata()
.upsert_task_execution_metadata(task_metadata_record)
.await?;
let key_count = context.data().len();
let keys: Vec<_> = context.data().keys().collect();
info!(
"Context saved: {} (workflow: {}, {} keys: {:?}, context_id: {:?})",
claimed_task.task_name, claimed_task.workflow_execution_id, key_count, keys, context_id
);
Ok(())
}
async fn should_retry_task(
&self,
claimed_task: &ClaimedTask,
error: &ExecutorError,
retry_policy: &RetryPolicy,
) -> Result<bool, ExecutorError> {
if matches!(error, ExecutorError::ClaimLost) {
return Ok(false);
}
if claimed_task.attempt >= retry_policy.max_attempts {
debug!(
"Task {} exceeded max retry attempts ({}/{})",
claimed_task.task_name, claimed_task.attempt, retry_policy.max_attempts
);
return Ok(false);
}
let should_retry = retry_policy
.retry_conditions
.iter()
.all(|condition| match condition {
RetryCondition::Never => false,
RetryCondition::AllErrors => true,
RetryCondition::TransientOnly => self.is_transient_error(error),
RetryCondition::ErrorPattern { patterns } => {
let error_msg = error.to_string().to_lowercase();
patterns
.iter()
.any(|pattern| error_msg.contains(&pattern.to_lowercase()))
}
});
debug!(
"Retry decision for task {}: {} (conditions: {:?}, error: {})",
claimed_task.task_name, should_retry, retry_policy.retry_conditions, error
);
Ok(should_retry)
}
pub fn is_transient_error(&self, error: &ExecutorError) -> bool {
match error {
ExecutorError::TaskTimeout => true,
ExecutorError::Database(_) => true,
ExecutorError::ConnectionPool(_) => true,
ExecutorError::TaskNotFound(_) => false,
ExecutorError::TaskExecution(task_error) => {
let error_msg = task_error.to_string().to_lowercase();
error_msg.contains("timeout")
|| error_msg.contains("connection")
|| error_msg.contains("network")
|| error_msg.contains("temporary")
|| error_msg.contains("unavailable")
}
_ => false,
}
}
async fn schedule_task_retry(
&self,
claimed_task: &ClaimedTask,
retry_policy: &RetryPolicy,
) -> Result<(), ExecutorError> {
let retry_delay = retry_policy.calculate_delay(claimed_task.attempt);
let retry_at = Utc::now() + retry_delay;
self.dal
.task_execution()
.schedule_retry(
claimed_task.task_execution_id,
crate::database::UniversalTimestamp(retry_at),
claimed_task.attempt + 1,
)
.await?;
info!(
"Scheduled retry for task {} in {:?} (attempt {})",
claimed_task.task_name,
retry_delay,
claimed_task.attempt + 1
);
Ok(())
}
}
#[cfg(all(test, feature = "sqlite"))]
mod is_transient_tests {
use super::*;
use crate::database::Database;
fn handler() -> TaskResultHandler {
let db = Database::new("sqlite://:memory:", "", 1);
let dal = DAL::new(db);
TaskResultHandler::new(
dal,
Arc::new(AtomicU64::new(0)),
Arc::new(AtomicU64::new(0)),
None,
)
}
#[test]
fn test_is_transient_timeout() {
assert!(handler().is_transient_error(&ExecutorError::TaskTimeout));
}
#[test]
fn test_is_transient_task_not_found() {
assert!(!handler().is_transient_error(&ExecutorError::TaskNotFound("missing".to_string())));
}
#[test]
fn test_is_transient_connection_pool() {
assert!(handler()
.is_transient_error(&ExecutorError::ConnectionPool("pool exhausted".to_string())));
}
#[test]
fn test_is_transient_task_execution_with_timeout_msg() {
let task_err = crate::error::TaskError::ExecutionFailed {
message: "connection timeout while waiting".to_string(),
task_id: "test".to_string(),
timestamp: chrono::Utc::now(),
};
assert!(handler().is_transient_error(&ExecutorError::TaskExecution(task_err)));
}
#[test]
fn test_is_transient_task_execution_permanent() {
let task_err = crate::error::TaskError::ExecutionFailed {
message: "invalid input data".to_string(),
task_id: "test".to_string(),
timestamp: chrono::Utc::now(),
};
assert!(!handler().is_transient_error(&ExecutorError::TaskExecution(task_err)));
}
#[test]
fn test_is_transient_task_execution_network() {
let task_err = crate::error::TaskError::ExecutionFailed {
message: "network unreachable".to_string(),
task_id: "test".to_string(),
timestamp: chrono::Utc::now(),
};
assert!(handler().is_transient_error(&ExecutorError::TaskExecution(task_err)));
}
#[test]
fn test_is_transient_task_execution_unavailable() {
let task_err = crate::error::TaskError::ExecutionFailed {
message: "service temporarily unavailable".to_string(),
task_id: "test".to_string(),
timestamp: chrono::Utc::now(),
};
assert!(handler().is_transient_error(&ExecutorError::TaskExecution(task_err)));
}
}