use super::error::LcelError;
use super::ext::RunnableExt;
use super::lambda::RunnableLambda;
use super::runnable_trait::Runnable;
use super::sequence::RunnableSequence;
use serde_json::Value;
use std::collections::HashMap;
pub trait RunnablePick<I: Send + Sync + 'static>: Sized {
fn pick<K>(
self,
keys: impl IntoIterator<Item = K>,
) -> RunnableSequence<I, HashMap<String, Value>>
where
K: Into<String>;
fn pluck(self, key: impl Into<String>) -> RunnableSequence<I, Value>;
}
impl<I, R> RunnablePick<I> for R
where
I: Send + Sync + 'static,
R: Runnable<I, HashMap<String, Value>> + Sized + 'static,
R::Error: Into<LcelError>,
{
fn pick<K>(
self,
keys: impl IntoIterator<Item = K>,
) -> RunnableSequence<I, HashMap<String, Value>>
where
K: Into<String>,
{
let keys: Vec<String> = keys.into_iter().map(Into::into).collect();
let filter = RunnableLambda::new_sync(move |m: HashMap<String, Value>| {
keys.iter()
.filter_map(|k| m.get(k).cloned().map(|v| (k.clone(), v)))
.collect::<HashMap<String, Value>>()
});
self.pipe(filter)
}
fn pluck(self, key: impl Into<String>) -> RunnableSequence<I, Value> {
let key = key.into();
let extract = RunnableLambda::new_sync(move |m: HashMap<String, Value>| {
m.get(&key).cloned().unwrap_or(Value::Null)
});
self.pipe(extract)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::Runnable;
use crate::RunnableLambda;
use crate::RunnableParallel;
fn parallel() -> RunnableParallel<String> {
RunnableParallel::<String>::new()
.with("len", RunnableLambda::new_sync(|s: String| s.len() as i64))
.with(
"upper",
RunnableLambda::new_sync(|s: String| s.to_uppercase()),
)
}
#[tokio::test]
async fn pick_keeps_only_selected_keys() {
let chain = parallel().pick(["len"]);
let out = chain.invoke("hello".to_string(), None).await.unwrap();
assert_eq!(out.len(), 1);
assert!(out.contains_key("len"));
assert!(!out.contains_key("upper"));
}
#[tokio::test]
async fn pick_multiple_keys_and_missing() {
let chain = parallel().pick(["len", "nope"]);
let out = chain.invoke("hello".to_string(), None).await.unwrap();
assert_eq!(out.len(), 1);
assert!(out.contains_key("len"));
}
#[tokio::test]
async fn pluck_returns_single_value() {
let chain = parallel().pluck("upper");
let out = chain.invoke("hello".to_string(), None).await.unwrap();
assert_eq!(out, Value::String("HELLO".to_string()));
}
#[tokio::test]
async fn pluck_missing_key_yields_null() {
let chain = parallel().pluck("missing");
let out = chain.invoke("hello".to_string(), None).await.unwrap();
assert_eq!(out, Value::Null);
}
#[tokio::test]
async fn pick_works_on_sequence_ending_in_map() {
let seq = RunnableLambda::new_sync(|s: String| s.to_uppercase())
.pipe(RunnableLambda::new_sync(|s: String| {
let mut m = HashMap::new();
m.insert("up".to_string(), Value::String(s));
m.insert("drop".to_string(), Value::Bool(true));
m
}))
.pick(["up"]);
let out = seq.invoke("hi".to_string(), None).await.unwrap();
assert_eq!(out.len(), 1);
assert_eq!(out["up"], Value::String("HI".to_string()));
}
}