moirai_iter/execution/
base.rs1use futures::StreamExt;
4use std::sync::Arc;
5
6use super::async_ctx::AsyncContext;
7use super::hybrid::HybridContext;
8use super::parallel::ParallelContext;
9
10const DEFAULT_ASYNC_CONCURRENCY: usize = 1024;
11
12pub trait ExecutionBase: Send + Sync {
14 fn context_type(&self) -> &'static str;
16
17 fn is_ready(&self) -> bool {
19 true
20 }
21}
22
23#[derive(Clone)]
26pub enum ExecutionContext {
27 Parallel(ParallelContext),
29 Async(AsyncContext),
31 Hybrid(HybridContext),
33}
34
35impl ExecutionContext {
36 pub fn execute<F, R>(&self, func: F) -> Result<R, Box<dyn std::error::Error + Send + Sync>>
38 where
39 F: FnOnce() -> R + Send,
40 R: Send,
41 {
42 match self {
43 ExecutionContext::Parallel(ctx) => ctx.execute(func),
44 ExecutionContext::Async(ctx) => ctx.execute(func),
45 ExecutionContext::Hybrid(ctx) => ctx.execute(func),
46 }
47 }
48
49 pub fn execute_iter<T, F, R>(
51 &self,
52 items: Vec<T>,
53 func: F,
54 ) -> Result<Vec<R>, Box<dyn std::error::Error + Send + Sync>>
55 where
56 T: Send + 'static,
57 F: Fn(T) -> R + Send + Sync + 'static,
58 R: Send + 'static,
59 {
60 match self {
61 ExecutionContext::Parallel(ctx) => ctx.execute_iter(items, func),
62 ExecutionContext::Async(ctx) => ctx.execute_iter(items, func),
63 ExecutionContext::Hybrid(ctx) => ctx.execute_iter(items, func),
64 }
65 }
66
67 pub async fn execute_async_iter<T, F, Fut, R>(
69 &self,
70 items: Vec<T>,
71 func: F,
72 ) -> Result<Vec<R>, Box<dyn std::error::Error + Send + Sync>>
73 where
74 T: Send + 'static,
75 F: Fn(T) -> Fut + Send + Sync + 'static,
76 Fut: std::future::Future<Output = R> + Send + 'static,
77 R: Send + 'static,
78 {
79 let concurrency = self.async_concurrency_limit();
80 let func = Arc::new(func);
81 let results = futures::stream::iter(items)
82 .map(|item| {
83 let func = Arc::clone(&func);
84 async move { func(item).await }
85 })
86 .buffered(concurrency)
87 .collect::<Vec<_>>()
88 .await;
89 Ok(results)
90 }
91
92 pub async fn execute_async_filter<T, F, Fut>(
94 &self,
95 items: Vec<T>,
96 predicate: F,
97 ) -> Result<Vec<T>, Box<dyn std::error::Error + Send + Sync>>
98 where
99 T: Send + 'static,
100 F: Fn(&T) -> Fut + Send + Sync + 'static,
101 Fut: std::future::Future<Output = bool> + Send + 'static,
102 {
103 let concurrency = self.async_concurrency_limit();
104 let predicate = Arc::new(predicate);
105 let results = futures::stream::iter(items)
106 .map(|item| {
107 let predicate = Arc::clone(&predicate);
108 async move {
109 let keep = predicate(&item).await;
110 (keep, item)
111 }
112 })
113 .buffered(concurrency)
114 .filter_map(|(keep, item)| async move { keep.then_some(item) })
115 .collect::<Vec<_>>()
116 .await;
117 Ok(results)
118 }
119
120 pub async fn execute_async_for_each<T, F, Fut>(
122 &self,
123 items: Vec<T>,
124 func: F,
125 ) -> Result<(), Box<dyn std::error::Error + Send + Sync>>
126 where
127 T: Send + 'static,
128 F: Fn(T) -> Fut + Send + Sync + 'static,
129 Fut: std::future::Future<Output = ()> + Send + 'static,
130 {
131 let concurrency = self.async_concurrency_limit();
132 let func = Arc::new(func);
133 futures::stream::iter(items)
134 .map(|item| {
135 let func = Arc::clone(&func);
136 async move { func(item).await }
137 })
138 .buffer_unordered(concurrency)
139 .collect::<Vec<_>>()
140 .await;
141 Ok(())
142 }
143
144 pub async fn execute_reduce<T, F>(
146 &self,
147 items: Vec<T>,
148 func: F,
149 ) -> Result<Option<T>, Box<dyn std::error::Error + Send + Sync>>
150 where
151 T: Send + 'static,
152 F: Fn(T, T) -> T + Send + Sync + 'static,
153 {
154 Ok(items.into_iter().reduce(func))
155 }
156
157 pub fn context_type(&self) -> &'static str {
159 match self {
160 ExecutionContext::Parallel(ctx) => ctx.context_type(),
161 ExecutionContext::Async(ctx) => ctx.context_type(),
162 ExecutionContext::Hybrid(ctx) => ctx.context_type(),
163 }
164 }
165
166 fn async_concurrency_limit(&self) -> usize {
167 match self {
168 ExecutionContext::Async(ctx) => ctx.max_concurrent,
169 ExecutionContext::Hybrid(ctx) => ctx.async_context.max_concurrent,
170 _ => std::thread::available_parallelism()
171 .map(|available| available.get())
172 .unwrap_or(DEFAULT_ASYNC_CONCURRENCY)
173 .max(1),
174 }
175 }
176}