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