use crate::command::NodeOutput;
use crate::config::GraphConfig;
use crate::error::AgentGraphError;
use crate::node::Node;
use crate::state::AgentState;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
pub type PayloadError = Box<dyn std::error::Error + Send + Sync>;
pub type TokenCallback = Arc<dyn Fn(&str) + Send + Sync>;
pub type InputSelector = Box<dyn Fn(&Value) -> Value + Send + Sync>;
pub type OutputMapper = Box<dyn Fn(&Value, &PayloadOutput) -> Value + Send + Sync>;
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct PayloadOutput {
pub value: Value,
#[serde(default)]
pub meta: HashMap<String, Value>,
}
pub struct PayloadContext {
pub on_token: Option<TokenCallback>,
pub run_id: String,
pub node_id: String,
}
pub trait Payload: Send + Sync {
fn invoke(
&self,
input: Value,
ctx: &PayloadContext,
) -> Pin<Box<dyn Future<Output = std::result::Result<PayloadOutput, PayloadError>> + Send + '_>>;
}
pub struct PayloadNode {
name: Option<String>,
payload: Box<dyn Payload>,
input_selector: Option<InputSelector>,
output_mapper: Option<OutputMapper>,
}
impl PayloadNode {
pub fn new(payload: Box<dyn Payload>) -> Self {
Self {
name: None,
payload,
input_selector: None,
output_mapper: None,
}
}
pub fn with_name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
pub fn with_input_selector(
mut self,
selector: impl Fn(&Value) -> Value + Send + Sync + 'static,
) -> Self {
self.input_selector = Some(Box::new(selector));
self
}
pub fn with_output_mapper(
mut self,
mapper: impl Fn(&Value, &PayloadOutput) -> Value + Send + Sync + 'static,
) -> Self {
self.output_mapper = Some(Box::new(mapper));
self
}
}
#[async_trait::async_trait]
impl Node for PayloadNode {
async fn execute(
&self,
state: &AgentState,
_config: &GraphConfig,
) -> crate::Result<NodeOutput> {
let state_data = state.export().await;
let state_value = serde_json::to_value(&state_data)
.map_err(|e| AgentGraphError::StateError(e.to_string()))?;
let input = match &self.input_selector {
Some(selector) => selector(&state_value),
None => state_value.clone(),
};
let ctx = PayloadContext {
on_token: None,
run_id: String::new(),
node_id: self.name.clone().unwrap_or_default(),
};
let output = self
.payload
.invoke(input, &ctx)
.await
.map_err(|e| AgentGraphError::PayloadError(e.to_string()))?;
let new_state_value = match &self.output_mapper {
Some(mapper) => mapper(&state_value, &output),
None => {
merge_value(&state_value, &output.value)
}
};
if let Value::Object(map) = new_state_value {
for (key, value) in map {
state.set_raw(&key, value).await?;
}
}
Ok(NodeOutput::Done)
}
fn name(&self) -> Option<&str> {
self.name.as_deref()
}
}
fn merge_value(state: &Value, output: &Value) -> Value {
match (state, output) {
(Value::Object(s), Value::Object(o)) => {
let mut result = s.clone();
for (k, v) in o {
result.insert(k.clone(), v.clone());
}
Value::Object(result)
}
_ => output.clone(),
}
}
impl std::fmt::Debug for PayloadNode {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PayloadNode")
.field("name", &self.name)
.finish()
}
}