pocketflow-async 0.1.0

Async runtime support for PocketFlow
Documentation
use pocketflow_core::*;
use std::collections::HashMap;
use std::sync::Arc;

type AsyncNodeFunc = Arc<
    dyn Fn(&mut (dyn std::any::Any + Send), &Params) -> Result<Option<String>> + Send + Sync,
>;

#[derive(Clone)]
pub struct AsyncNode {
    name: String,
    params: Params,
    successors: HashMap<String, AsyncNode>,
    func: AsyncNodeFunc,
}

impl AsyncNode {
    pub fn new<F>(name: impl Into<String>, func: F) -> Self
    where
        F: Fn(&mut (dyn std::any::Any + Send), &Params) -> Result<Option<String>>
            + Send
            + Sync
            + 'static,
    {
        Self {
            name: name.into(),
            params: Params::new(),
            successors: HashMap::new(),
            func: Arc::new(func),
        }
    }

    pub fn add_successor(&mut self, action: impl Into<String>, node: AsyncNode) -> &mut Self {
        self.successors.insert(action.into(), node);
        self
    }

    pub fn next(&mut self, node: AsyncNode) -> &mut Self {
        self.add_successor("default", node)
    }

    pub fn set_params(&mut self, params: Params) {
        self.params = params;
    }

    pub fn get_params(&self) -> &Params {
        &self.params
    }

    pub fn get_successor(&self, action: &str) -> Option<&AsyncNode> {
        self.successors.get(action)
    }

    pub fn has_successors(&self) -> bool {
        !self.successors.is_empty()
    }

    pub async fn run(&self, shared: &mut (dyn std::any::Any + Send)) -> Result<()> {
        if self.has_successors() {
            eprintln!("Warning: AsyncNode won't run successors. Use AsyncFlow.");
        }
        
        (self.func)(shared, &self.params)?;
        Ok(())
    }

    pub async fn run_recursive(&self, shared: &mut (dyn std::any::Any + Send)) -> Result<()> {
        let action = (self.func)(shared, &self.params)?;
        
        if let Some(next_node) = action
            .as_ref()
            .and_then(|a| self.successors.get(a))
            .or_else(|| self.successors.get("default")) {
            Box::pin(next_node.run_recursive(shared)).await?;
        }
        
        Ok(())
    }
}

#[derive(Clone)]
pub struct AsyncFlow {
    start_node: Option<AsyncNode>,
    params: Params,
}

impl AsyncFlow {
    pub fn new() -> Self {
        Self {
            start_node: None,
            params: Params::new(),
        }
    }

    pub fn start(mut self, node: AsyncNode) -> Self {
        self.start_node = Some(node);
        self
    }

    pub fn set_params(&mut self, params: Params) {
        self.params = params;
    }

    pub async fn run(&self, shared: &mut (dyn std::any::Any + Send)) -> Result<()> {
        if let Some(ref node) = self.start_node {
            let mut node = node.clone();
            node.set_params(self.params.clone());
            node.run_recursive(shared).await?;
        }
        Ok(())
    }

    pub async fn run_with_params(&self, shared: &mut (dyn std::any::Any + Send), params: Params) -> Result<()> {
        if let Some(ref node) = self.start_node {
            let mut node = node.clone();
            let mut merged_params = self.params.clone();
            merged_params.merge(&params);
            node.set_params(merged_params);
            node.run_recursive(shared).await?;
        }
        Ok(())
    }
}

impl Default for AsyncFlow {
    fn default() -> Self {
        Self::new()
    }
}

#[derive(Clone)]
pub struct AsyncBatchFlow {
    start_node: Option<AsyncNode>,
    params: Params,
}

impl AsyncBatchFlow {
    pub fn new() -> Self {
        Self {
            start_node: None,
            params: Params::new(),
        }
    }

    pub fn start(mut self, node: AsyncNode) -> Self {
        self.start_node = Some(node);
        self
    }

    pub fn set_params(&mut self, params: Params) {
        self.params = params;
    }

    pub async fn run_batch(&self, shared: &mut (dyn std::any::Any + Send), batch_params: Vec<Params>) -> Result<Vec<()>> {
        let mut results = Vec::with_capacity(batch_params.len());
        for params in batch_params {
            let flow = AsyncFlow {
                start_node: self.start_node.clone(),
                params: self.params.clone(),
            };
            flow.run_with_params(shared, params).await?;
            results.push(());
        }
        Ok(results)
    }
}

impl Default for AsyncBatchFlow {
    fn default() -> Self {
        Self::new()
    }
}

#[cfg(test)]
mod tests {
    use super::*;
    use std::sync::{Arc, Mutex};

    #[derive(Default, Clone)]
    struct TestShared {
        pub counter: Arc<Mutex<i32>>,
    }

    #[tokio::test]
    async fn test_async_flow_execution() {
        let mut shared = TestShared::default();
        
        let node = AsyncNode::new("test", |shared, _params| {
            if let Some(shared) = shared.downcast_mut::<TestShared>() {
                let mut counter = shared.counter.lock().unwrap();
                *counter += 1;
            }
            Ok(None)
        });
        
        let flow = AsyncFlow::new().start(node);
        
        flow.run(&mut shared).await.unwrap();
        
        let counter = shared.counter.lock().unwrap();
        assert_eq!(*counter, 1);
    }

    #[tokio::test]
    async fn test_async_chained_flow() {
        let mut shared = TestShared::default();
        
        let mut node1 = AsyncNode::new("node1", |shared, _params| {
            if let Some(shared) = shared.downcast_mut::<TestShared>() {
                let mut counter = shared.counter.lock().unwrap();
                *counter += 1;
            }
            Ok(None)
        });
        
        let node2 = AsyncNode::new("node2", |shared, _params| {
            if let Some(shared) = shared.downcast_mut::<TestShared>() {
                let mut counter = shared.counter.lock().unwrap();
                *counter += 10;
            }
            Ok(None)
        });
        
        node1.next(node2);
        let flow = AsyncFlow::new().start(node1);
        
        flow.run(&mut shared).await.unwrap();
        
        let counter = shared.counter.lock().unwrap();
        assert_eq!(*counter, 11);
    }
}