use crate::core::context_data::ContextData; use crate::core::control::PipelineControl;
use crate::error::{OrkaError, OrkaResult};
use std::any::Any;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
pub type Handler<TData, Err> = Box<
dyn Fn(ContextData<TData>) -> Pin<Box<dyn Future<Output = Result<PipelineControl, Err>> + Send>>
+ Send
+ Sync,
>;
pub trait AnyContextDataExtractor<TData: 'static + Send + Sync>: Send + Sync {
fn extract_sub_context_data(&self, root_ctx_data: ContextData<TData>) -> OrkaResult<Box<dyn Any + Send>>;
fn sub_context_data_type_id(&self) -> std::any::TypeId;
fn merge_sub_context_data(
&self,
_root_ctx_data: ContextData<TData>,
_sub_ctx_data: &(dyn Any + Send),
) -> OrkaResult<()> {
Ok(())
}
fn has_merge(&self) -> bool {
false
}
}
pub type MergeFn<TData, SData> = Arc<dyn Fn(&mut TData, &SData) + Send + Sync + 'static>;
pub type ExtractorFn<TData, SData> =
Arc<dyn Fn(ContextData<TData>) -> Result<ContextData<SData>, OrkaError> + Send + Sync + 'static>;
pub type ConditionFn<TData> = Arc<dyn Fn(ContextData<TData>) -> bool + Send + Sync + 'static>;
pub struct ContextDataExtractorImpl<
TData: 'static + Send + Sync,
SData: 'static + Send + Sync, > {
extractor_fn: Arc<dyn Fn(ContextData<TData>) -> OrkaResult<ContextData<SData>> + Send + Sync + 'static>,
merge_fn: Option<MergeFn<TData, SData>>,
}
impl<TData: 'static + Send + Sync, SData: 'static + Send + Sync> ContextDataExtractorImpl<TData, SData> {
pub fn new(f: impl Fn(ContextData<TData>) -> OrkaResult<ContextData<SData>> + Send + Sync + 'static) -> Self {
Self {
extractor_fn: Arc::new(f),
merge_fn: None,
}
}
pub fn with_merge(
f: impl Fn(ContextData<TData>) -> OrkaResult<ContextData<SData>> + Send + Sync + 'static,
merge: impl Fn(&mut TData, &SData) + Send + Sync + 'static,
) -> Self {
Self {
extractor_fn: Arc::new(f),
merge_fn: Some(Arc::new(merge)),
}
}
}
impl<TData: 'static + Send + Sync, SData: 'static + Send + Sync> AnyContextDataExtractor<TData>
for ContextDataExtractorImpl<TData, SData>
{
fn extract_sub_context_data(&self, root_ctx_data: ContextData<TData>) -> OrkaResult<Box<dyn Any + Send>> {
let sub_ctx_data: ContextData<SData> = (self.extractor_fn)(root_ctx_data)?;
Ok(Box::new(sub_ctx_data))
}
fn sub_context_data_type_id(&self) -> std::any::TypeId {
std::any::TypeId::of::<SData>()
}
fn merge_sub_context_data(
&self,
root_ctx_data: ContextData<TData>,
sub_ctx_data: &(dyn Any + Send),
) -> OrkaResult<()> {
let Some(merge) = self.merge_fn.as_ref() else {
return Ok(());
};
let sub: &ContextData<SData> = sub_ctx_data.downcast_ref::<ContextData<SData>>().ok_or_else(|| {
OrkaError::Internal(format!(
"Internal type mismatch merging sub-context back: expected ContextData<{}>.",
std::any::type_name::<SData>()
))
})?;
let sub_guard = sub.read();
let mut root_guard = root_ctx_data.write();
merge(&mut *root_guard, &*sub_guard);
Ok(())
}
fn has_merge(&self) -> bool {
self.merge_fn.is_some()
}
}
pub(crate) fn downcast_context_data<SData: 'static + Send + Sync>(
any_ctx_data: Box<dyn Any + Send>,
expected_sdata_type_id: std::any::TypeId, step_name: &str,
) -> OrkaResult<ContextData<SData>> {
if std::any::TypeId::of::<SData>() != expected_sdata_type_id {
return Err(OrkaError::TypeMismatch {
step_name: step_name.to_string(),
expected_type: format!(
"ContextData<{}> (underlying SData TypeId: {:?})",
std::any::type_name::<SData>(),
std::any::TypeId::of::<SData>()
),
});
}
match any_ctx_data.downcast::<ContextData<SData>>() {
Ok(boxed_ctx_data) => Ok(*boxed_ctx_data), Err(_) => {
Err(OrkaError::Internal(format!(
"Internal type mismatch during ContextData downcast for step '{}'. Expected ContextData<{}> but downcast failed despite TypeId match.",
step_name,
std::any::type_name::<SData>()
)))
}
}
}