use crate::core::context_data::ContextData;
use crate::error::OrkaError;
use crate::pipeline::Pipeline;
use async_trait::async_trait;
use std::future::Future;
use std::marker::PhantomData;
use std::sync::Arc;
#[async_trait]
pub trait PipelineProvider<TData, SData, MainErr>: Send + Sync + 'static
where
TData: 'static + Send + Sync,
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
async fn get_pipeline(&self, main_ctx_data: ContextData<TData>) -> Result<Arc<Pipeline<SData, MainErr>>, OrkaError>;
}
#[derive(Clone)]
pub struct DynPipelineProvider<TData, SData, MainErr>
where
TData: 'static + Send + Sync,
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
inner: Arc<dyn PipelineProvider<TData, SData, MainErr>>,
}
impl<TData, SData, MainErr> DynPipelineProvider<TData, SData, MainErr>
where
TData: 'static + Send + Sync,
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
pub fn new(inner: Arc<dyn PipelineProvider<TData, SData, MainErr>>) -> Self {
Self { inner }
}
}
#[async_trait]
impl<TData, SData, MainErr> PipelineProvider<TData, SData, MainErr> for DynPipelineProvider<TData, SData, MainErr>
where
TData: 'static + Send + Sync,
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
async fn get_pipeline(&self, main_ctx_data: ContextData<TData>) -> Result<Arc<Pipeline<SData, MainErr>>, OrkaError> {
self.inner.get_pipeline(main_ctx_data).await
}
}
#[derive(Clone)]
pub struct StaticPipelineProvider<SData, MainErr>
where
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
pipeline: Arc<Pipeline<SData, MainErr>>,
_phantom_sdata: PhantomData<SData>,
_phantom_main_err: PhantomData<MainErr>,
}
impl<SData, MainErr> StaticPipelineProvider<SData, MainErr>
where
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
pub fn new(pipeline: Arc<Pipeline<SData, MainErr>>) -> Self {
Self {
pipeline,
_phantom_sdata: PhantomData,
_phantom_main_err: PhantomData,
}
}
}
#[async_trait]
impl<TData, SData, MainErr> PipelineProvider<TData, SData, MainErr> for StaticPipelineProvider<SData, MainErr>
where
TData: 'static + Send + Sync,
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
async fn get_pipeline(&self, _main_ctx_data: ContextData<TData>) -> Result<Arc<Pipeline<SData, MainErr>>, OrkaError> {
Ok(self.pipeline.clone())
}
}
pub struct FunctionalPipelineProvider<TData, SData, MainErr, F, Fut>
where
TData: 'static + Send + Sync,
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
F: Fn(ContextData<TData>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Arc<Pipeline<SData, MainErr>>, OrkaError>> + Send + 'static,
{
factory: F,
_phantom_tdata: PhantomData<fn() -> TData>,
_phantom_sdata: PhantomData<fn() -> SData>,
_phantom_main_err: PhantomData<fn() -> MainErr>,
}
impl<TData, SData, MainErr, F, Fut> FunctionalPipelineProvider<TData, SData, MainErr, F, Fut>
where
TData: 'static + Send + Sync,
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
F: Fn(ContextData<TData>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Arc<Pipeline<SData, MainErr>>, OrkaError>> + Send + 'static,
{
pub fn new(factory: F) -> Self {
Self {
factory,
_phantom_tdata: PhantomData,
_phantom_sdata: PhantomData,
_phantom_main_err: PhantomData,
}
}
}
#[async_trait]
impl<TData, SData, MainErr, F, Fut> PipelineProvider<TData, SData, MainErr>
for FunctionalPipelineProvider<TData, SData, MainErr, F, Fut>
where
TData: 'static + Send + Sync,
SData: 'static + Send + Sync,
MainErr: std::error::Error + From<OrkaError> + Send + Sync + 'static,
F: Fn(ContextData<TData>) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<Arc<Pipeline<SData, MainErr>>, OrkaError>> + Send + 'static,
{
async fn get_pipeline(&self, main_ctx_data: ContextData<TData>) -> Result<Arc<Pipeline<SData, MainErr>>, OrkaError> {
let sdata_type_name = std::any::type_name::<SData>();
(self.factory)(main_ctx_data).await.map_err(|orka_err_from_factory| {
OrkaError::PipelineProviderFailure {
step_name: format!("functional_provider_for_{}", sdata_type_name),
source: anyhow::anyhow!(
"Factory for SData='{}' failed: {}",
sdata_type_name,
orka_err_from_factory
),
}
})
}
}