use async_trait::async_trait;
use mofa_kernel::workflow::{Reducer, ReducerType, StateUpdate};
use serde_json::Value;
use mofa_kernel::agent::error::{AgentError, AgentResult};
#[derive(Debug, Clone, Default)]
pub struct OverwriteReducer;
#[async_trait]
impl Reducer for OverwriteReducer {
async fn reduce(&self, _current: Option<&Value>, update: &Value) -> AgentResult<Value> {
Ok(update.clone())
}
fn name(&self) -> &str {
"overwrite"
}
fn reducer_type(&self) -> ReducerType {
ReducerType::Overwrite
}
}
#[derive(Debug, Clone, Default)]
pub struct AppendReducer;
#[async_trait]
impl Reducer for AppendReducer {
async fn reduce(&self, current: Option<&Value>, update: &Value) -> AgentResult<Value> {
let mut arr = match current {
Some(Value::Array(a)) => a.clone(),
_ => Vec::new(),
};
arr.push(update.clone());
Ok(Value::Array(arr))
}
fn name(&self) -> &str {
"append"
}
fn reducer_type(&self) -> ReducerType {
ReducerType::Append
}
}
#[derive(Debug, Clone, Default)]
pub struct ExtendReducer;
#[async_trait]
impl Reducer for ExtendReducer {
async fn reduce(&self, current: Option<&Value>, update: &Value) -> AgentResult<Value> {
let mut arr = match current {
Some(Value::Array(a)) => a.clone(),
_ => Vec::new(),
};
match update {
Value::Array(items) => {
arr.extend(items.iter().cloned());
}
other => {
arr.push(other.clone());
}
}
Ok(Value::Array(arr))
}
fn name(&self) -> &str {
"extend"
}
fn reducer_type(&self) -> ReducerType {
ReducerType::Extend
}
}
#[derive(Debug, Clone)]
pub struct MergeReducer {
pub deep: bool,
}
impl Default for MergeReducer {
fn default() -> Self {
Self { deep: false }
}
}
impl MergeReducer {
pub fn shallow() -> Self {
Self { deep: false }
}
pub fn deep() -> Self {
Self { deep: true }
}
}
#[async_trait]
impl Reducer for MergeReducer {
async fn reduce(&self, current: Option<&Value>, update: &Value) -> AgentResult<Value> {
match (current, update) {
(Some(Value::Object(current_map)), Value::Object(update_map)) => {
let mut result = current_map.clone();
for (key, value) in update_map {
if self.deep {
if let (Some(Value::Object(existing)), Value::Object(new_obj)) =
(result.get(key), value)
{
let merged = merge_objects_deep(existing.clone(), new_obj.clone());
result.insert(key.clone(), Value::Object(merged));
continue;
}
}
result.insert(key.clone(), value.clone());
}
Ok(Value::Object(result))
}
(None, Value::Object(update_map)) => Ok(Value::Object(update_map.clone())),
(Some(current), _) => Ok(current.clone()),
(None, update) => Ok(update.clone()),
}
}
fn name(&self) -> &str {
if self.deep { "merge_deep" } else { "merge" }
}
fn reducer_type(&self) -> ReducerType {
ReducerType::Merge { deep: self.deep }
}
}
fn merge_objects_deep(
mut base: serde_json::Map<String, Value>,
update: serde_json::Map<String, Value>,
) -> serde_json::Map<String, Value> {
for (key, value) in update {
match (base.get(&key), value) {
(Some(Value::Object(base_obj)), Value::Object(update_obj)) => {
let merged = merge_objects_deep(base_obj.clone(), update_obj);
base.insert(key, Value::Object(merged));
}
(_, value) => {
base.insert(key, value);
}
}
}
base
}
#[derive(Debug, Clone)]
pub struct LastNReducer {
pub n: usize,
}
impl LastNReducer {
pub fn new(n: usize) -> Self {
Self { n }
}
}
#[async_trait]
impl Reducer for LastNReducer {
async fn reduce(&self, current: Option<&Value>, update: &Value) -> AgentResult<Value> {
let mut arr = match current {
Some(Value::Array(a)) => a.clone(),
_ => Vec::new(),
};
match update {
Value::Array(items) => {
arr.extend(items.iter().cloned());
}
other => {
arr.push(other.clone());
}
}
if arr.len() > self.n {
let start = arr.len() - self.n;
arr = arr.split_off(start);
}
Ok(Value::Array(arr))
}
fn name(&self) -> &str {
"last_n"
}
fn reducer_type(&self) -> ReducerType {
ReducerType::LastN { n: self.n }
}
}
#[derive(Debug, Clone, Default)]
pub struct FirstReducer;
#[async_trait]
impl Reducer for FirstReducer {
async fn reduce(&self, current: Option<&Value>, update: &Value) -> AgentResult<Value> {
match current {
Some(value) if !value.is_null() => Ok(value.clone()),
_ => Ok(update.clone()),
}
}
fn name(&self) -> &str {
"first"
}
fn reducer_type(&self) -> ReducerType {
ReducerType::First
}
}
#[derive(Debug, Clone, Default)]
pub struct LastReducer;
#[async_trait]
impl Reducer for LastReducer {
async fn reduce(&self, _current: Option<&Value>, update: &Value) -> AgentResult<Value> {
if update.is_null() {
Ok(update.clone())
} else {
Ok(update.clone())
}
}
fn name(&self) -> &str {
"last"
}
fn reducer_type(&self) -> ReducerType {
ReducerType::Last
}
}
pub struct CustomReducer<F>
where
F: Fn(Option<&Value>, &Value) -> AgentResult<Value> + Send + Sync,
{
name: String,
func: F,
}
impl<F> CustomReducer<F>
where
F: Fn(Option<&Value>, &Value) -> AgentResult<Value> + Send + Sync,
{
pub fn new(name: impl Into<String>, func: F) -> Self {
Self {
name: name.into(),
func,
}
}
}
#[async_trait]
impl<F> Reducer for CustomReducer<F>
where
F: Fn(Option<&Value>, &Value) -> AgentResult<Value> + Send + Sync,
{
async fn reduce(&self, current: Option<&Value>, update: &Value) -> AgentResult<Value> {
(self.func)(current, update)
}
fn name(&self) -> &str {
&self.name
}
fn reducer_type(&self) -> ReducerType {
ReducerType::Custom(self.name.clone())
}
}
pub fn create_reducer(reducer_type: &ReducerType) -> AgentResult<Box<dyn Reducer>> {
match reducer_type {
ReducerType::Overwrite => Ok(Box::new(OverwriteReducer)),
ReducerType::Append => Ok(Box::new(AppendReducer)),
ReducerType::Extend => Ok(Box::new(ExtendReducer)),
ReducerType::Merge { deep } => Ok(Box::new(MergeReducer { deep: *deep })),
ReducerType::LastN { n } => Ok(Box::new(LastNReducer::new(*n))),
ReducerType::First => Ok(Box::new(FirstReducer)),
ReducerType::Last => Ok(Box::new(LastReducer)),
ReducerType::Custom(name) => Err(AgentError::Internal(format!(
"Cannot create reducer for unknown custom type: {}",
name
))),
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[tokio::test]
async fn test_overwrite_reducer() {
let reducer = OverwriteReducer;
let result = reducer
.reduce(Some(&json!("old")), &json!("new"))
.await
.unwrap();
assert_eq!(result, json!("new"));
let result = reducer.reduce(None, &json!("value")).await.unwrap();
assert_eq!(result, json!("value"));
}
#[tokio::test]
async fn test_append_reducer() {
let reducer = AppendReducer;
let result = reducer
.reduce(Some(&json!(["a", "b"])), &json!("c"))
.await
.unwrap();
assert_eq!(result, json!(["a", "b", "c"]));
let result = reducer.reduce(None, &json!("first")).await.unwrap();
assert_eq!(result, json!(["first"]));
}
#[tokio::test]
async fn test_extend_reducer() {
let reducer = ExtendReducer;
let result = reducer
.reduce(Some(&json!([1, 2])), &json!([3, 4]))
.await
.unwrap();
assert_eq!(result, json!([1, 2, 3, 4]));
let result = reducer.reduce(Some(&json!([1])), &json!(2)).await.unwrap();
assert_eq!(result, json!([1, 2]));
}
#[tokio::test]
async fn test_merge_reducer_shallow() {
let reducer = MergeReducer::shallow();
let current = json!({"a": 1, "b": 2});
let update = json!({"b": 3, "c": 4});
let result = reducer.reduce(Some(¤t), &update).await.unwrap();
assert_eq!(result["a"], 1);
assert_eq!(result["b"], 3);
assert_eq!(result["c"], 4);
}
#[tokio::test]
async fn test_merge_reducer_deep() {
let reducer = MergeReducer::deep();
let current = json!({
"config": {
"a": 1,
"b": { "x": 1, "y": 2 }
}
});
let update = json!({
"config": {
"b": { "y": 3, "z": 4 },
"c": 5
}
});
let result = reducer.reduce(Some(¤t), &update).await.unwrap();
assert_eq!(result["config"]["a"], 1);
assert_eq!(result["config"]["b"]["x"], 1);
assert_eq!(result["config"]["b"]["y"], 3);
assert_eq!(result["config"]["b"]["z"], 4);
assert_eq!(result["config"]["c"], 5);
}
#[tokio::test]
async fn test_last_n_reducer() {
let reducer = LastNReducer::new(3);
let result = reducer
.reduce(Some(&json!([1, 2, 3, 4])), &json!(5))
.await
.unwrap();
assert_eq!(result, json!([3, 4, 5]));
let result = reducer
.reduce(Some(&json!([1, 2])), &json!([3, 4, 5, 6]))
.await
.unwrap();
assert_eq!(result, json!([4, 5, 6]));
}
#[tokio::test]
async fn test_first_reducer() {
let reducer = FirstReducer;
let result = reducer.reduce(None, &json!("first")).await.unwrap();
assert_eq!(result, json!("first"));
let result = reducer
.reduce(Some(&json!("first")), &json!("second"))
.await
.unwrap();
assert_eq!(result, json!("first"));
let result = reducer
.reduce(Some(&json!(null)), &json!("value"))
.await
.unwrap();
assert_eq!(result, json!("value"));
}
#[tokio::test]
async fn test_last_reducer() {
let reducer = LastReducer;
let result = reducer
.reduce(Some(&json!("first")), &json!("second"))
.await
.unwrap();
assert_eq!(result, json!("second"));
let result = reducer
.reduce(Some(&json!("old")), &json!("new"))
.await
.unwrap();
assert_eq!(result, json!("new"));
}
#[tokio::test]
async fn test_custom_reducer() {
let reducer = CustomReducer::new("sum", |current, update| {
let curr = current.and_then(|v| v.as_i64()).unwrap_or(0);
let upd = update.as_i64().unwrap_or(0);
Ok(json!(curr + upd))
});
let result = reducer.reduce(Some(&json!(10)), &json!(5)).await.unwrap();
assert_eq!(result, json!(15));
assert_eq!(reducer.name(), "sum");
}
#[test]
fn test_create_reducer() {
let r = create_reducer(&ReducerType::Overwrite).unwrap();
assert_eq!(r.name(), "overwrite");
let r = create_reducer(&ReducerType::Append).unwrap();
assert_eq!(r.name(), "append");
let r = create_reducer(&ReducerType::LastN { n: 5 }).unwrap();
assert_eq!(r.name(), "last_n");
}
}