use std::{
pin::Pin,
task::{Context, Poll},
};
use apalis_codec::json::JsonCodec;
use futures_channel::mpsc::SendError;
use futures_core::{Stream, stream::BoxStream};
use futures_util::{StreamExt, TryStreamExt, stream};
use serde::{Serialize, de::Deserialize};
use serde_json::Value;
use apalis_core::{
backend::{Backend, BackendExt, TaskStream, queue::Queue},
task::{Task, status::Status, task_id::RandomId},
worker::{context::WorkerContext, ext::ack::AcknowledgeLayer},
};
use crate::{
JsonMapMetadata, JsonStorage,
util::{FindFirstWith, JsonAck},
};
impl<Args> Backend for JsonStorage<Args>
where
Args: 'static + Send + Serialize + for<'de> Deserialize<'de> + Unpin,
{
type Args = Args;
type IdType = RandomId;
type Error = SendError;
type Context = JsonMapMetadata;
type Stream = TaskStream<Task<Args, JsonMapMetadata, RandomId>, SendError>;
type Layer = AcknowledgeLayer<JsonAck<Args>>;
type Beat = BoxStream<'static, Result<(), Self::Error>>;
fn heartbeat(&self, _: &WorkerContext) -> Self::Beat {
stream::once(async { Ok(()) }).boxed()
}
fn middleware(&self) -> Self::Layer {
AcknowledgeLayer::new(JsonAck {
inner: self.clone(),
})
}
fn poll(self, _worker: &WorkerContext) -> Self::Stream {
(self.map(|r| Ok(Some(r))).boxed()) as _
}
}
impl<Args: 'static + Send + Serialize + for<'de> Deserialize<'de> + Unpin> BackendExt
for JsonStorage<Args>
{
type Codec = JsonCodec<Value>;
type Compact = Value;
type CompactStream = TaskStream<Task<Self::Compact, JsonMapMetadata, RandomId>, SendError>;
fn get_queue(&self) -> Queue {
std::any::type_name::<Args>().into()
}
fn poll_compact(self, worker: &WorkerContext) -> Self::CompactStream {
self.poll(worker)
.map_ok(|c| {
c.map(|t| t.map(|args| serde_json::to_value(args).expect("to be encodable")))
})
.boxed()
}
}
impl<Args: for<'de> Deserialize<'de> + Unpin> Stream for JsonStorage<Args> {
type Item = Task<Args, JsonMapMetadata, RandomId>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let map = self.tasks.try_write().unwrap();
if let Some((key, task)) = map.find_first_with(|s, _| {
s.queue == std::any::type_name::<Args>() && s.status == Status::Pending
}) {
use apalis_core::task::builder::TaskBuilder;
let key = key.clone();
let args = Args::deserialize(&task.args).unwrap();
let task = TaskBuilder::new(args)
.with_task_id(key.task_id.clone())
.with_ctx(task.ctx.clone())
.build();
drop(map);
let this = self.get_mut();
this.update_status(&key, Status::Running)
.expect("Failed to update status");
this.persist_to_disk().expect("Failed to persist to disk");
Poll::Ready(Some(task))
} else {
Poll::Pending
}
}
}