use std::{
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use futures_channel::mpsc::SendError;
use futures_core::{Stream, stream::BoxStream};
use futures_sink::Sink;
use futures_util::{SinkExt, StreamExt};
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use serde_json::Value;
use apalis_core::{
backend::memory::{MemorySink, MemoryStorage},
task::{
Task,
status::Status,
task_id::{RandomId, TaskId},
},
};
use crate::{
JsonMapMetadata, JsonStorage,
util::{FindFirstWith, TaskKey, TaskWithMeta},
};
#[derive(Debug)]
struct SharedJsonStream<T, Ctx> {
inner: JsonStorage<Value>,
req_type: std::marker::PhantomData<(T, Ctx)>,
}
impl<Args: DeserializeOwned + Unpin> Stream for SharedJsonStream<Args, JsonMapMetadata> {
type Item = Task<Args, JsonMapMetadata, RandomId>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
use apalis_core::task::builder::TaskBuilder;
let map = self.inner.tasks.try_read().expect("Failed to read tasks");
if let Some((key, _)) = map.find_first_with(|k, _| {
k.queue == std::any::type_name::<Args>() && k.status == Status::Pending
}) {
let task = map.get(key).unwrap();
let Ok(args) = Args::deserialize(&task.args) else {
return Poll::Pending;
};
let task = TaskBuilder::new(args)
.with_task_id(key.task_id.clone())
.with_ctx(task.ctx.clone())
.build();
let key = key.clone();
drop(map);
let this = &mut self.get_mut().inner;
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
}
}
}
#[derive(Debug, Clone)]
pub struct SharedJsonStore {
inner: JsonStorage<serde_json::Value>,
}
impl Default for SharedJsonStore {
fn default() -> Self {
Self::new()
}
}
impl SharedJsonStore {
#[must_use]
pub fn new() -> Self {
Self {
inner: JsonStorage::new_temp().unwrap(),
}
}
}
impl<Args: Send + Serialize + for<'de> Deserialize<'de> + Unpin + 'static>
apalis_core::backend::shared::MakeShared<Args> for SharedJsonStore
{
type Backend = MemoryStorage<Args, JsonMapMetadata>;
type Config = ();
type MakeError = String;
fn make_shared_with_config(
&mut self,
_: Self::Config,
) -> Result<Self::Backend, Self::MakeError> {
let (sender, receiver) = self.inner.create_channel::<Args>();
let sender = MemorySink::new(Arc::new(futures_util::lock::Mutex::new(sender)));
Ok(MemoryStorage::new_with(sender, receiver))
}
}
type BoxSink<Args> = Box<
dyn Sink<Task<Args, JsonMapMetadata, RandomId>, Error = SendError>
+ Send
+ Sync
+ Unpin
+ 'static,
>;
impl JsonStorage<Value> {
fn create_channel<Args: 'static + for<'de> Deserialize<'de> + Serialize + Send + Unpin>(
&self,
) -> (
BoxSink<Args>,
BoxStream<'static, Task<Args, JsonMapMetadata, RandomId>>,
) {
let sender = self.clone();
let wrapped_sender = {
let store = self.clone();
sender.with_flat_map(move |task: Task<Args, JsonMapMetadata, RandomId>| {
use apalis_core::task::task_id::RandomId;
let task_id = task
.parts
.task_id
.clone()
.unwrap_or(TaskId::new(RandomId::default()));
let task = task.map(|args| serde_json::to_value(args).unwrap());
store
.insert(
&TaskKey {
task_id,
queue: std::any::type_name::<Args>().to_owned(),
status: Status::Pending,
},
TaskWithMeta {
args: task.args.clone(),
ctx: task.parts.ctx.clone(),
result: None,
idempotency_key: task.parts.idempotency_key.clone(),
},
)
.unwrap();
futures_util::stream::iter(vec![Ok(task)])
})
};
let filtered_stream = {
let inner = self.clone();
SharedJsonStream {
inner,
req_type: std::marker::PhantomData,
}
};
let sender = Box::new(wrapped_sender)
as Box<
dyn Sink<Task<Args, JsonMapMetadata, RandomId>, Error = SendError>
+ Send
+ Sync
+ Unpin,
>;
let receiver = filtered_stream.boxed();
(sender, receiver)
}
}
#[cfg(test)]
mod tests {
use std::time::Duration;
use apalis_core::error::BoxDynError;
use apalis_core::worker::context::WorkerContext;
use apalis_core::{
backend::{TaskSink, shared::MakeShared},
worker::{builder::WorkerBuilder, ext::event_listener::EventListenerExt},
};
use super::*;
const ITEMS: u32 = 10;
#[tokio::test]
async fn basic_shared() {
let mut store = SharedJsonStore::new();
let mut string_store = store.make_shared().unwrap();
let mut int_store = store.make_shared().unwrap();
for i in 0..ITEMS {
string_store.push(format!("ITEM: {i}")).await.unwrap();
int_store.push(i).await.unwrap();
}
async fn task(task: u32, ctx: WorkerContext) -> Result<(), BoxDynError> {
tokio::time::sleep(Duration::from_millis(2)).await;
if task == ITEMS - 1 {
ctx.stop()?;
return Err("Worker stopped!")?;
}
Ok(())
}
let string_worker = WorkerBuilder::new("rango-tango-string")
.backend(string_store)
.on_event(|ctx, ev| {
println!("CTX {:?}, On Event = {ev:?}", ctx.name());
})
.build(|req: String, ctx: WorkerContext| async move {
tokio::time::sleep(Duration::from_millis(2)).await;
println!("{req}");
if req.ends_with(&(ITEMS - 1).to_string()) {
ctx.stop().unwrap();
}
})
.run();
let int_worker = WorkerBuilder::new("rango-tango-int")
.backend(int_store)
.on_event(|ctx, ev| {
println!("CTX {:?}, On Event = {ev:?}", ctx.name());
})
.build(task)
.run();
let _ = futures_util::future::join(int_worker, string_worker).await;
}
}