use crate::core::control::PipelineResult;
use crate::core::context_data::ContextData;
use crate::core::trace::RunOutcome;
use crate::error::OrkaError;
use crate::pipeline::definition::Pipeline as CorePipeline;
use crate::error::OrkaResult;
use crate::pipeline::runner::PipelineRunner;
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>;
async fn run_any_erased_detailed(
&self,
ctx_obj: Box<dyn Any + Send>,
) -> (Result<PipelineResult, ApplicationError>, RunOutcome);
fn as_any(&self) -> &dyn Any;
}
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,
{
runner: Arc<dyn PipelineRunner<TData, PipelineHandlerError>>,
pipeline: Option<Arc<CorePipeline<TData, 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.runner.run(typed_ctx_data).await.map_err(ApplicationError::from)
}
async fn run_any_erased_detailed(
&self,
ctx_obj: Box<dyn Any + Send>,
) -> (Result<PipelineResult, ApplicationError>, RunOutcome) {
let typed_ctx_data = match ctx_obj.downcast::<ContextData<TData>>() {
Ok(boxed_ctx_data) => *boxed_ctx_data,
Err(_) => {
let orka_type_mismatch = OrkaError::TypeMismatch {
step_name: "registry_dispatch".to_string(),
expected_type: std::any::type_name::<ContextData<TData>>().to_string(),
};
let outcome = RunOutcome::Errored {
step: "registry_dispatch".to_string(),
message: orka_type_mismatch.to_string(),
};
return (Err(ApplicationError::from(orka_type_mismatch)), outcome);
}
};
let (result, outcome) = self.runner.run_with_outcome(typed_ctx_data).await;
(result.map_err(ApplicationError::from), outcome)
}
fn as_any(&self) -> &dyn Any {
self
}
}
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 pipeline = Arc::new(pipeline);
let wrapper = PipelineWrapper::<TData, PipelineHandlerError, ApplicationError> {
runner: pipeline.clone(),
pipeline: Some(pipeline),
_phantom_app_err: PhantomData,
};
self
.registry
.lock()
.insert(TypeId::of::<TData>(), Arc::new(wrapper));
Ok(())
}
pub fn register_runner<TData, PipelineHandlerError>(
&self,
runner: Arc<dyn PipelineRunner<TData, PipelineHandlerError>>,
) 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 runner.");
let wrapper = PipelineWrapper::<TData, PipelineHandlerError, ApplicationError> {
runner,
pipeline: None,
_phantom_app_err: PhantomData,
};
self
.registry
.lock()
.insert(TypeId::of::<TData>(), Arc::new(wrapper));
}
pub fn pipeline<TData, PipelineHandlerError>(&self) -> Option<Arc<CorePipeline<TData, PipelineHandlerError>>>
where
TData: 'static + Send + Sync,
PipelineHandlerError: std::error::Error + From<OrkaError> + Send + Sync + 'static,
ApplicationError: From<PipelineHandlerError>,
CorePipeline<TData, PipelineHandlerError>: Send + Sync,
{
let registry = self.registry.lock();
let runner = registry.get(&TypeId::of::<TData>())?;
let wrapper = runner
.as_any()
.downcast_ref::<PipelineWrapper<TData, PipelineHandlerError, ApplicationError>>()?;
wrapper.pipeline.clone()
}
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
}
pub async fn run_with_outcome<TData>(
&self,
ctx_data: ContextData<TData>,
) -> (Result<PipelineResult, ApplicationError>, RunOutcome)
where
TData: 'static + Send + Sync,
{
let type_id = TypeId::of::<TData>();
let runner_arc = { self.registry.lock().get(&type_id).cloned() };
let Some(runner_arc) = runner_arc else {
let type_name = std::any::type_name::<TData>();
let orka_config_err = OrkaError::ConfigurationError {
step_name: "Orka::run".to_string(),
message: format!("No pipeline registered for TData type {}", type_name),
};
let outcome = RunOutcome::Errored {
step: "Orka::run".to_string(),
message: orka_config_err.to_string(),
};
return (Err(ApplicationError::from(orka_config_err)), outcome);
};
let owned_ctx_obj: Box<dyn Any + Send> = Box::new(ctx_data.clone());
runner_arc.run_any_erased_detailed(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()
}
}