Skip to main content

moirai_iter/execution/
base.rs

1//! Base trait and enum for execution contexts.
2
3use 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
11/// Base trait for all execution contexts
12pub trait ExecutionBase: Send + Sync {
13    /// Get context type name for debugging
14    fn context_type(&self) -> &'static str;
15
16    /// Check if the context is ready for execution
17    fn is_ready(&self) -> bool {
18        true
19    }
20}
21
22/// Concrete execution context enum that wraps different strategy implementations
23/// This approach ensures type safety while avoiding dyn-compatibility issues
24#[derive(Clone)]
25pub enum ExecutionContext {
26    /// Parallel execution for CPU-bound work
27    Parallel(ParallelContext),
28    /// Async execution for I/O-bound work  
29    Async(AsyncContext),
30    /// Hybrid execution that adapts between strategies
31    Hybrid(HybridContext),
32}
33
34impl ExecutionContext {
35    /// Execute a function once with the appropriate context
36    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    /// Execute an iterator operation with proper type erasure
49    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    /// Execute async iterator operations
67    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    /// Execute async filter operations
86    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    /// Execute async for_each operations
110    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    /// Execute parallel reduce operations
128    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    /// Get context type name
141    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}