use futures::StreamExt;
use std::sync::Arc;
use super::async_ctx::AsyncContext;
use super::hybrid::HybridContext;
use super::parallel::ParallelContext;
const DEFAULT_ASYNC_CONCURRENCY: usize = 1024;
pub trait ExecutionBase: Send + Sync {
fn context_type(&self) -> &'static str;
fn is_ready(&self) -> bool {
true
}
}
#[derive(Clone)]
pub enum ExecutionContext {
Parallel(ParallelContext),
Async(AsyncContext),
Hybrid(HybridContext),
}
impl ExecutionContext {
pub fn execute<F, R>(&self, func: F) -> Result<R, Box<dyn std::error::Error + Send + Sync>>
where
F: FnOnce() -> R + Send,
R: Send,
{
match self {
ExecutionContext::Parallel(ctx) => ctx.execute(func),
ExecutionContext::Async(ctx) => ctx.execute(func),
ExecutionContext::Hybrid(ctx) => ctx.execute(func),
}
}
pub fn execute_iter<T, F, R>(
&self,
items: Vec<T>,
func: F,
) -> Result<Vec<R>, Box<dyn std::error::Error + Send + Sync>>
where
T: Send + 'static,
F: Fn(T) -> R + Send + Sync + 'static,
R: Send + 'static,
{
match self {
ExecutionContext::Parallel(ctx) => ctx.execute_iter(items, func),
ExecutionContext::Async(ctx) => ctx.execute_iter(items, func),
ExecutionContext::Hybrid(ctx) => ctx.execute_iter(items, func),
}
}
pub async fn execute_async_iter<T, F, Fut, R>(
&self,
items: Vec<T>,
func: F,
) -> Result<Vec<R>, Box<dyn std::error::Error + Send + Sync>>
where
T: Send + 'static,
F: Fn(T) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = R> + Send + 'static,
R: Send + 'static,
{
let concurrency = self.async_concurrency_limit();
let func = Arc::new(func);
let results = futures::stream::iter(items)
.map(|item| {
let func = Arc::clone(&func);
async move { func(item).await }
})
.buffered(concurrency)
.collect::<Vec<_>>()
.await;
Ok(results)
}
pub async fn execute_async_filter<T, F, Fut>(
&self,
items: Vec<T>,
predicate: F,
) -> Result<Vec<T>, Box<dyn std::error::Error + Send + Sync>>
where
T: Send + 'static,
F: Fn(&T) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = bool> + Send + 'static,
{
let concurrency = self.async_concurrency_limit();
let predicate = Arc::new(predicate);
let results = futures::stream::iter(items)
.map(|item| {
let predicate = Arc::clone(&predicate);
async move {
let keep = predicate(&item).await;
(keep, item)
}
})
.buffered(concurrency)
.filter_map(|(keep, item)| async move { keep.then_some(item) })
.collect::<Vec<_>>()
.await;
Ok(results)
}
pub async fn execute_async_for_each<T, F, Fut>(
&self,
items: Vec<T>,
func: F,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>>
where
T: Send + 'static,
F: Fn(T) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
let concurrency = self.async_concurrency_limit();
let func = Arc::new(func);
futures::stream::iter(items)
.map(|item| {
let func = Arc::clone(&func);
async move { func(item).await }
})
.buffer_unordered(concurrency)
.collect::<Vec<_>>()
.await;
Ok(())
}
pub async fn execute_reduce<T, F>(
&self,
items: Vec<T>,
func: F,
) -> Result<Option<T>, Box<dyn std::error::Error + Send + Sync>>
where
T: Send + 'static,
F: Fn(T, T) -> T + Send + Sync + 'static,
{
Ok(items.into_iter().reduce(func))
}
pub fn context_type(&self) -> &'static str {
match self {
ExecutionContext::Parallel(ctx) => ctx.context_type(),
ExecutionContext::Async(ctx) => ctx.context_type(),
ExecutionContext::Hybrid(ctx) => ctx.context_type(),
}
}
fn async_concurrency_limit(&self) -> usize {
match self {
ExecutionContext::Async(ctx) => ctx.max_concurrent,
ExecutionContext::Hybrid(ctx) => ctx.async_context.max_concurrent,
_ => std::thread::available_parallelism()
.map(|available| available.get())
.unwrap_or(DEFAULT_ASYNC_CONCURRENCY)
.max(1),
}
}
}