Skip to main content

moirai_iter/execution/
base.rs

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