use crate::error::{BeaverError, BeaverResult};
use crate::fixed_count_task::FixedCountTask;
use crate::periodic_task::PeriodicTask;
use crate::platform;
use crate::range_interval_task::RangeIntervalTask;
use crate::task::Task;
use crate::time_interval_task::TimeIntervalTask;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::sync::Mutex;
use std::time::Duration;
use tokio::runtime::Handle;
use tokio::sync::mpsc;
use tokio::task::JoinHandle;
use tokio::time::sleep;
struct Splash {
task: Arc<Task>,
seq: u64,
}
pub(crate) struct Dam {
tx: Mutex<Option<mpsc::Sender<Splash>>>,
current: Arc<Mutex<Option<Arc<Task>>>>,
release_flag: AtomicBool,
enqueue_seq: AtomicU64,
cancel_watermark: Arc<AtomicU64>,
worker: Mutex<Option<JoinHandle<()>>>,
}
impl Dam {
#[inline]
pub(crate) fn new(name: impl Into<String>, buffer: usize) -> Self {
Self::with_capacity(name, buffer)
}
pub(crate) fn with_capacity(_name: impl Into<String>, buffer: usize) -> Self {
let current = Arc::new(Mutex::new(None));
let current_worker = Arc::clone(¤t);
let cancel_watermark = Arc::new(AtomicU64::new(0));
let watermark_worker = Arc::clone(&cancel_watermark);
let (tx, mut rx) = mpsc::channel::<Splash>(buffer);
let join = platform::spawn(async move {
while let Some(msg) = rx.recv().await {
run_loop_msg(¤t_worker, &watermark_worker, msg).await;
}
});
Self {
tx: Mutex::new(Some(tx)),
current,
release_flag: AtomicBool::new(false),
enqueue_seq: AtomicU64::new(0),
cancel_watermark,
worker: Mutex::new(Some(join)),
}
}
pub(crate) fn with_handle(_name: impl Into<String>, capacity: usize, handle: Handle) -> Self {
let current = Arc::new(Mutex::new(None));
let current_worker = Arc::clone(¤t);
let cancel_watermark = Arc::new(AtomicU64::new(0));
let watermark_worker = Arc::clone(&cancel_watermark);
let (tx, mut rx) = mpsc::channel::<Splash>(capacity);
let join = platform::spawn_on(&handle, async move {
while let Some(msg) = rx.recv().await {
run_loop_msg(¤t_worker, &watermark_worker, msg).await;
}
});
Self {
tx: Mutex::new(Some(tx)),
current,
release_flag: AtomicBool::new(false),
enqueue_seq: AtomicU64::new(0),
cancel_watermark,
worker: Mutex::new(Some(join)),
}
}
pub(crate) async fn enqueue(&self, task: Arc<Task>) -> BeaverResult<()> {
if self.release_flag.load(Ordering::Acquire) {
return Err(BeaverError::DamReleased);
}
let seq = self.enqueue_seq.fetch_add(1, Ordering::AcqRel);
let guard = self.tx.lock()?;
match guard.as_ref() {
Some(tx) => tx
.try_send(Splash { task, seq })
.map_err(|_| BeaverError::QueueFull),
None => Err(BeaverError::DamReleased),
}
}
pub(crate) async fn cancel_all(&self) -> BeaverResult<()> {
let watermark = self.enqueue_seq.load(Ordering::Acquire);
self.cancel_watermark.fetch_max(watermark, Ordering::AcqRel);
if let Some(s) = self.current.lock()?.as_ref() {
s.set_interrupted(true);
}
Ok(())
}
pub(crate) async fn release(&self) -> BeaverResult<()> {
self.release_flag.store(true, Ordering::Release);
self.cancel_watermark.store(u64::MAX, Ordering::Release);
if let Some(s) = self.current.lock()?.as_ref() {
s.set_interrupted(true);
}
let _ = self.tx.lock()?.take();
Ok(())
}
pub(crate) fn take_worker(&self) -> Option<JoinHandle<()>> {
self.worker.lock().ok().and_then(|mut g| g.take())
}
}
impl Drop for Dam {
fn drop(&mut self) {
self.release_flag.store(true, Ordering::Release);
self.cancel_watermark.store(u64::MAX, Ordering::Release);
if let Ok(guard) = self.current.lock() {
if let Some(s) = guard.as_ref() {
s.set_interrupted(true);
}
}
}
}
const INTERRUPT_ORDERING: Ordering = Ordering::Relaxed;
#[inline]
async fn run_time_interval(task: &TimeIntervalTask) {
let intervals = &task.intervals[..];
let work = &task.work;
let listener = task.listener.as_ref();
for (i, &millis) in intervals.iter().enumerate() {
if task.interrupted.load(INTERRUPT_ORDERING) {
if let Some(l) = listener {
l.on_interrupt();
}
return;
}
if millis > 0 {
sleep(Duration::from_millis(millis)).await;
}
let result = work.execute().await;
if !result.need_retry() {
if let Some(l) = listener {
l.on_complete();
}
return;
}
if i == intervals.len() - 1 {
if let Some(l) = listener {
l.on_error(crate::error::RuntimeError::RetriesExhausted);
}
}
}
}
#[inline]
async fn run_range_interval(task: &RangeIntervalTask) {
let total = task.total_retries as usize;
let intervals = &task.intervals[..];
let work = &task.work;
let listener = task.listener.as_ref();
for attempt in 0..total {
if task.interrupted.load(INTERRUPT_ORDERING) {
if let Some(l) = listener {
l.on_interrupt();
}
return;
}
if attempt > 0 {
let millis = intervals[attempt - 1];
if millis > 0 {
sleep(Duration::from_millis(millis)).await;
}
}
let result = work.execute().await;
if !result.need_retry() {
if let Some(l) = listener {
l.on_complete();
}
return;
}
if attempt == total - 1 {
if let Some(l) = listener {
l.on_error(crate::error::RuntimeError::RetriesExhausted);
}
}
}
}
#[inline]
async fn run_fixed_count(task: &FixedCountTask) {
let total = task.count;
let work = &task.work;
let progress = task.progress.as_ref();
let listener = task.listener.as_ref();
let tag = task.tag.as_deref().map_or("", |v| v);
for current in 1..=total {
if task.interrupted.load(INTERRUPT_ORDERING) {
if let Some(l) = listener {
l.on_interrupt();
}
return;
}
if let Some(p) = progress {
p.on_progress(current, total, tag);
}
let result = work.execute().await;
if !result.need_retry() {
if let Some(l) = listener {
l.on_complete();
}
return;
}
if current == total {
if let Some(l) = listener {
l.on_error(crate::error::RuntimeError::RetriesExhausted);
}
}
}
}
#[inline]
async fn run_periodic(task: &PeriodicTask) {
let interval = task.interval;
let work = &task.work;
let listener = task.listener.as_ref();
if task.initial_delay && !interval.is_zero() {
sleep(interval).await;
}
loop {
if task.interrupted.load(INTERRUPT_ORDERING) {
if let Some(l) = listener {
l.on_interrupt();
}
return;
}
let result = work.execute().await;
if !result.need_retry() {
if let Some(l) = listener {
l.on_complete();
}
return;
}
if !interval.is_zero() {
sleep(interval).await;
}
}
}
#[inline]
async fn run_task(task: &Task) {
match task {
Task::TimeInterval(s) => run_time_interval(s).await,
Task::RangeInterval(s) => run_range_interval(s).await,
Task::FixedCount(s) => run_fixed_count(s).await,
Task::Periodic(s) => run_periodic(s).await,
}
}
fn panic_message_to_string(payload: Box<dyn std::any::Any + Send>) -> String {
if let Some(s) = payload.downcast_ref::<&'static str>() {
return (*s).to_string();
}
if let Ok(s) = payload.downcast::<String>() {
return *s;
}
"panic (unknown payload)".to_string()
}
fn notify_error(task: &Task, error: crate::error::RuntimeError) {
match task {
Task::TimeInterval(t) => {
if let Some(l) = &t.listener {
l.on_error(error);
}
}
Task::RangeInterval(t) => {
if let Some(l) = &t.listener {
l.on_error(error);
}
}
Task::FixedCount(t) => {
if let Some(l) = &t.listener {
l.on_error(error);
}
}
Task::Periodic(t) => {
if let Some(l) = &t.listener {
l.on_error(error);
}
}
}
}
async fn run_loop_msg(
current_worker: &Mutex<Option<Arc<Task>>>,
cancel_watermark: &AtomicU64,
splash: Splash,
) {
let Splash { task, seq } = splash;
if seq < cancel_watermark.load(Ordering::Acquire) {
task.interrupt();
return;
}
{
let mut guard = current_worker.lock().expect("mutex poisoned");
*guard = Some(Arc::clone(&task));
}
loop {
let task_for_join = Arc::clone(&task);
let join = platform::spawn(async move { run_task(task_for_join.as_ref()).await });
match join.await {
Ok(()) => break,
Err(join_err) => {
if !join_err.is_panic() {
break;
}
let msg = panic_message_to_string(join_err.into_panic());
notify_error(&task, crate::error::RuntimeError::TaskExecutionFailed(msg));
let restart_interval = match task.as_ref() {
Task::Periodic(p) if !task.interrupted() => Some(p.interval),
_ => None,
};
match restart_interval {
Some(interval) => {
if !interval.is_zero() {
sleep(interval).await;
}
}
None => break,
}
}
}
}
{
let mut guard = current_worker.lock().expect("mutex poisoned");
*guard = None;
}
}