lc_core/runnables/
pick.rs1use super::error::LcelError;
18use super::ext::RunnableExt;
19use super::lambda::RunnableLambda;
20use super::runnable_trait::Runnable;
21use super::sequence::RunnableSequence;
22use serde_json::Value;
23use std::collections::HashMap;
24
25pub trait RunnablePick<I: Send + Sync + 'static>: Sized {
27 fn pick<K>(self, keys: impl IntoIterator<Item = K>) -> RunnableSequence<I, HashMap<String, Value>>
41 where
42 K: Into<String>;
43
44 fn pluck(self, key: impl Into<String>) -> RunnableSequence<I, Value>;
55}
56
57impl<I, R> RunnablePick<I> for R
58where
59 I: Send + Sync + 'static,
60 R: Runnable<I, HashMap<String, Value>> + Sized + 'static,
61 R::Error: Into<LcelError>,
62{
63 fn pick<K>(self, keys: impl IntoIterator<Item = K>) -> RunnableSequence<I, HashMap<String, Value>>
64 where
65 K: Into<String>,
66 {
67 let keys: Vec<String> = keys.into_iter().map(Into::into).collect();
68 let filter = RunnableLambda::new_sync(move |m: HashMap<String, Value>| {
69 keys.iter()
70 .filter_map(|k| m.get(k).cloned().map(|v| (k.clone(), v)))
71 .collect::<HashMap<String, Value>>()
72 });
73 self.pipe(filter)
74 }
75
76 fn pluck(self, key: impl Into<String>) -> RunnableSequence<I, Value> {
77 let key = key.into();
78 let extract = RunnableLambda::new_sync(move |m: HashMap<String, Value>| {
79 m.get(&key).cloned().unwrap_or(Value::Null)
80 });
81 self.pipe(extract)
82 }
83}
84
85#[cfg(test)]
86mod tests {
87 use super::*;
88 use crate::Runnable;
89 use crate::RunnableLambda;
90 use crate::RunnableParallel;
91
92 fn parallel() -> RunnableParallel<String> {
93 RunnableParallel::<String>::new()
94 .with("len", RunnableLambda::new_sync(|s: String| s.len() as i64))
95 .with("upper", RunnableLambda::new_sync(|s: String| s.to_uppercase()))
96 }
97
98 #[tokio::test]
99 async fn pick_keeps_only_selected_keys() {
100 let chain = parallel().pick(["len"]);
101 let out = chain.invoke("hello".to_string(), None).await.unwrap();
102 assert_eq!(out.len(), 1);
103 assert!(out.contains_key("len"));
104 assert!(!out.contains_key("upper"));
105 }
106
107 #[tokio::test]
108 async fn pick_multiple_keys_and_missing() {
109 let chain = parallel().pick(["len", "nope"]);
110 let out = chain.invoke("hello".to_string(), None).await.unwrap();
111 assert_eq!(out.len(), 1);
113 assert!(out.contains_key("len"));
114 }
115
116 #[tokio::test]
117 async fn pluck_returns_single_value() {
118 let chain = parallel().pluck("upper");
119 let out = chain.invoke("hello".to_string(), None).await.unwrap();
120 assert_eq!(out, Value::String("HELLO".to_string()));
121 }
122
123 #[tokio::test]
124 async fn pluck_missing_key_yields_null() {
125 let chain = parallel().pluck("missing");
126 let out = chain.invoke("hello".to_string(), None).await.unwrap();
127 assert_eq!(out, Value::Null);
128 }
129
130 #[tokio::test]
131 async fn pick_works_on_sequence_ending_in_map() {
132 let seq = RunnableLambda::new_sync(|s: String| s.to_uppercase())
134 .pipe(RunnableLambda::new_sync(|s: String| {
135 let mut m = HashMap::new();
136 m.insert("up".to_string(), Value::String(s));
137 m.insert("drop".to_string(), Value::Bool(true));
138 m
139 }))
140 .pick(["up"]);
141 let out = seq.invoke("hi".to_string(), None).await.unwrap();
142 assert_eq!(out.len(), 1);
143 assert_eq!(out["up"], Value::String("HI".to_string()));
144 }
145}