use crate::core::control::PipelineResult;
use crate::core::context_data::ContextData;
use crate::error::OrkaError; use crate::pipeline::definition::Pipeline as CorePipeline;
use crate::error::OrkaResult;
use async_trait::async_trait;
use parking_lot::Mutex;
use std::any::{Any, TypeId};
use std::collections::HashMap;
use std::marker::PhantomData;
use std::sync::Arc;
use tracing::{event, instrument, Level};
#[async_trait]
trait AnyPipelineRunner<ApplicationError>: Send + Sync
where
ApplicationError: std::error::Error + Send + Sync + 'static,
{
async fn run_any_erased_with_owned_ctx(&self, ctx_obj: Box<dyn Any + Send>) -> Result<PipelineResult, ApplicationError>;
}
struct PipelineWrapper<TData, PipelineHandlerError, ApplicationError>
where
TData: 'static + Send + Sync,
PipelineHandlerError: std::error::Error + From<OrkaError> + Send + Sync + 'static, ApplicationError: std::error::Error + From<PipelineHandlerError> + From<OrkaError> + Send + Sync + 'static,
CorePipeline<TData, PipelineHandlerError>: Send + Sync,
{
pipeline: Arc<CorePipeline<TData, PipelineHandlerError>>,
_phantom_tdata: PhantomData<TData>,
_phantom_handler_err: PhantomData<PipelineHandlerError>,
_phantom_app_err: PhantomData<ApplicationError>,
}
#[async_trait]
impl<TData, PipelineHandlerError, ApplicationError> AnyPipelineRunner<ApplicationError>
for PipelineWrapper<TData, PipelineHandlerError, ApplicationError>
where
TData: 'static + Send + Sync,
PipelineHandlerError: std::error::Error + From<OrkaError> + Send + Sync + 'static,
ApplicationError: std::error::Error + From<PipelineHandlerError> + From<OrkaError> + Send + Sync + 'static,
CorePipeline<TData, PipelineHandlerError>: Send + Sync,
{
#[instrument(
name = "PipelineWrapper::run_any_erased_with_owned_ctx",
skip_all,
fields(
target_tdata_type = %std::any::type_name::<TData>(),
pipeline_handler_error_type = %std::any::type_name::<PipelineHandlerError>(),
application_error_type = %std::any::type_name::<ApplicationError>(),
),
err(Display)
)]
async fn run_any_erased_with_owned_ctx(&self, ctx_obj: Box<dyn Any + Send>) -> Result<PipelineResult, ApplicationError> {
event!(Level::TRACE, "Attempting to downcast owned context object.");
let typed_ctx_data = match ctx_obj.downcast::<ContextData<TData>>() {
Ok(boxed_ctx_data) => *boxed_ctx_data, Err(_) => { let expected_type_name = std::any::type_name::<ContextData<TData>>();
event!(Level::ERROR, "Context object type mismatch. Expected {}.", expected_type_name);
let orka_type_mismatch = OrkaError::TypeMismatch {
step_name: "registry_dispatch".to_string(),
expected_type: expected_type_name.to_string(),
};
return Err(ApplicationError::from(orka_type_mismatch));
}
};
event!(Level::DEBUG, "Context object downcast successful. Executing wrapped pipeline.");
self.pipeline.run(typed_ctx_data).await.map_err(ApplicationError::from)
}
}
pub struct Orka<ApplicationError = OrkaError>
where
ApplicationError: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
registry: Mutex<HashMap<TypeId, Arc<dyn AnyPipelineRunner<ApplicationError>>>>,
_phantom_app_err: PhantomData<ApplicationError>,
}
impl<ApplicationError> Orka<ApplicationError>
where
ApplicationError: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
pub fn new() -> Self {
Self {
registry: Mutex::new(HashMap::new()),
_phantom_app_err: PhantomData,
}
}
pub fn register_pipeline<TData, PipelineHandlerError>(
&self,
pipeline: CorePipeline<TData, PipelineHandlerError>,
) -> OrkaResult<()>
where
TData: 'static + Send + Sync,
PipelineHandlerError: std::error::Error + From<OrkaError> + Send + Sync + 'static, ApplicationError: From<PipelineHandlerError>, CorePipeline<TData, PipelineHandlerError>: Send + Sync,
{
event!(Level::DEBUG, tdata_type = %std::any::type_name::<TData>(), pipeline_handler_error = %std::any::type_name::<PipelineHandlerError>(), "Registering pipeline.");
pipeline.validate()?;
let wrapper = PipelineWrapper::<TData, PipelineHandlerError, ApplicationError> {
pipeline: Arc::new(pipeline),
_phantom_tdata: PhantomData,
_phantom_handler_err: PhantomData,
_phantom_app_err: PhantomData,
};
self
.registry
.lock()
.insert(TypeId::of::<TData>(), Arc::new(wrapper));
Ok(())
}
pub async fn run<TData>(&self, ctx_data: ContextData<TData>) -> Result<PipelineResult, ApplicationError>
where
TData: 'static + Send + Sync,
{
event!(Level::DEBUG, tdata_type = %std::any::type_name::<TData>(), "Attempting to run pipeline.");
let type_id = TypeId::of::<TData>();
let runner_arc: Arc<dyn AnyPipelineRunner<ApplicationError>>;
{
let reg_lock = self.registry.lock();
runner_arc = reg_lock
.get(&type_id)
.cloned()
.ok_or_else(|| {
let type_name = std::any::type_name::<TData>();
event!(Level::ERROR, "No pipeline registered for TData type {}.", type_name);
let orka_config_err = OrkaError::ConfigurationError {
step_name: "Orka::run".to_string(),
message: format!("No pipeline registered for TData type {}", type_name),
};
ApplicationError::from(orka_config_err)
})?;
}
let owned_ctx_obj: Box<dyn Any + Send> = Box::new(ctx_data.clone());
runner_arc.run_any_erased_with_owned_ctx(owned_ctx_obj).await
}
}
impl<ApplicationError> Default for Orka<ApplicationError>
where
ApplicationError: std::error::Error + From<OrkaError> + Send + Sync + 'static,
{
fn default() -> Self {
Self::new()
}
}
impl Orka<OrkaError> {
pub fn new_default() -> Self {
Orka::<OrkaError>::new()
}
}