1use std::collections::BTreeMap;
10use std::fmt;
11use std::future::Future;
12use std::pin::Pin;
13use std::sync::Arc;
14
15use serde::Deserialize;
16use serde::de::DeserializeOwned;
17use serde_json::Value;
18use taquba_workflow::{Step, StepError};
19
20use crate::subprocess::Subprocess;
21use crate::task::TaskIdentity;
22
23#[derive(Debug, Clone, Copy)]
25pub struct Task<'a> {
26 pub step: &'a Step,
29 pub identity: &'a TaskIdentity,
31 pub params: &'a Value,
33 pub inputs: &'a BTreeMap<String, Value>,
35}
36
37#[derive(Debug, Clone, PartialEq, Eq)]
40pub enum Outcome {
41 Succeeded(Value),
43 Failed(String),
45}
46
47pub trait Operator: Send + Sync + 'static {
49 type Params: DeserializeOwned + Send;
52
53 fn run(
56 &self,
57 task: &Task<'_>,
58 params: Self::Params,
59 ) -> impl Future<Output = Result<Outcome, StepError>> + Send;
60}
61
62type Check = Box<dyn Fn(&toml::Table) -> Result<(), String> + Send + Sync>;
63type BoxFuture<'a> = Pin<Box<dyn Future<Output = Result<Outcome, StepError>> + Send + 'a>>;
64type Run = Box<dyn for<'a> Fn(&'a Task<'a>, Value) -> BoxFuture<'a> + Send + Sync>;
65
66struct Entry {
67 check: Check,
68 run: Option<Run>,
69}
70
71#[derive(Default)]
73pub struct OperatorSet {
74 entries: BTreeMap<String, Entry>,
75}
76
77#[derive(Debug, Clone, PartialEq, Eq)]
79pub enum OperatorError {
80 Unknown,
82 InvalidParams(String),
85}
86
87impl OperatorSet {
88 pub fn new() -> Self {
90 Self::default()
91 }
92
93 pub fn builtin() -> Self {
95 let mut set = Self::new();
96 set.add("subprocess", Subprocess::default());
97 set
98 }
99
100 pub fn register<P: DeserializeOwned + 'static>(&mut self, name: &str) {
103 self.entries.insert(
104 name.to_string(),
105 Entry {
106 check: check_fn::<P>(),
107 run: None,
108 },
109 );
110 }
111
112 pub fn add<O: Operator>(&mut self, name: &str, operator: O) {
115 let operator = Arc::new(operator);
116 let run: Run = Box::new(move |task, params| {
117 let operator = operator.clone();
118 Box::pin(async move {
119 let params: O::Params = serde_json::from_value(params).map_err(|e| {
120 StepError::permanent(format!("the rendered parameters are invalid: {e}"))
121 })?;
122 operator.run(task, params).await
123 })
124 });
125 self.entries.insert(
126 name.to_string(),
127 Entry {
128 check: check_fn::<O::Params>(),
129 run: Some(run),
130 },
131 );
132 }
133
134 pub fn contains(&self, name: &str) -> bool {
136 self.entries.contains_key(name)
137 }
138
139 pub fn runs(&self, name: &str) -> bool {
141 self.entries.get(name).is_some_and(|e| e.run.is_some())
142 }
143
144 pub fn names(&self) -> impl Iterator<Item = &str> {
146 self.entries.keys().map(String::as_str)
147 }
148
149 pub fn check(&self, name: &str, params: &toml::Table) -> Result<(), OperatorError> {
151 let entry = self.entries.get(name).ok_or(OperatorError::Unknown)?;
152 (entry.check)(params).map_err(OperatorError::InvalidParams)
153 }
154
155 pub async fn run(
158 &self,
159 name: &str,
160 task: &Task<'_>,
161 params: Value,
162 ) -> Result<Outcome, StepError> {
163 let run = self
164 .entries
165 .get(name)
166 .and_then(|e| e.run.as_ref())
167 .ok_or_else(|| {
168 StepError::permanent(format!("operator `{name}` cannot run in this process"))
169 })?;
170 run(task, params).await
171 }
172}
173
174fn check_fn<P: DeserializeOwned + 'static>() -> Check {
175 Box::new(|params: &toml::Table| {
176 P::deserialize(params.clone())
177 .map(|_| ())
178 .map_err(|e| e.to_string())
179 })
180}
181
182impl fmt::Debug for OperatorSet {
183 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
184 f.debug_set().entries(self.entries.keys()).finish()
185 }
186}
187
188#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
190#[serde(deny_unknown_fields)]
191pub struct SubprocessParams {
192 #[serde(deserialize_with = "non_empty_argv")]
194 pub argv: Vec<String>,
195}
196
197fn non_empty_argv<'de, D: serde::Deserializer<'de>>(
198 deserializer: D,
199) -> Result<Vec<String>, D::Error> {
200 let argv = Vec::<String>::deserialize(deserializer)?;
201 if argv.is_empty() {
202 return Err(serde::de::Error::custom("argv must not be empty"));
203 }
204 Ok(argv)
205}
206
207#[cfg(test)]
208mod tests {
209 use super::*;
210 use crate::partition::Partition;
211
212 fn table(text: &str) -> toml::Table {
213 text.parse().unwrap()
214 }
215
216 #[test]
217 fn builtin_set_runs_the_subprocess_operator() {
218 let set = OperatorSet::builtin();
219 assert_eq!(set.names().collect::<Vec<_>>(), ["subprocess"]);
220 assert!(set.runs("subprocess"));
221 assert_eq!(
222 set.check("subprocess", &table(r#"argv = ["python", "x.py"]"#)),
223 Ok(())
224 );
225 }
226
227 #[test]
228 fn unknown_operator_is_reported_as_unknown() {
229 assert_eq!(
230 OperatorSet::builtin().check("sql", &table("")),
231 Err(OperatorError::Unknown)
232 );
233 }
234
235 #[test]
236 fn params_that_do_not_deserialize_are_invalid() {
237 let set = OperatorSet::builtin();
238 for params in ["argv = []", "argv = [\"a\"]\nextra = 1", ""] {
239 assert!(
240 matches!(
241 set.check("subprocess", &table(params)),
242 Err(OperatorError::InvalidParams(_))
243 ),
244 "{params}"
245 );
246 }
247 }
248
249 #[test]
250 fn register_adds_a_check_only_operator() {
251 #[derive(Deserialize)]
252 struct Params {
253 #[allow(dead_code)]
254 query: String,
255 }
256 let mut set = OperatorSet::new();
257 assert!(!set.contains("sql"));
258 set.register::<Params>("sql");
259 assert!(set.contains("sql"));
260 assert!(!set.runs("sql"));
261 assert_eq!(set.check("sql", &table(r#"query = "select 1""#)), Ok(()));
262 }
263
264 struct Echo;
265
266 #[derive(Deserialize)]
267 struct EchoParams {
268 text: String,
269 }
270
271 impl Operator for Echo {
272 type Params = EchoParams;
273
274 async fn run(&self, task: &Task<'_>, params: EchoParams) -> Result<Outcome, StepError> {
275 Ok(Outcome::Succeeded(serde_json::json!({
276 "text": params.text,
277 "node": task.identity.node,
278 })))
279 }
280 }
281
282 fn task<'a>(
283 step: &'a Step,
284 identity: &'a TaskIdentity,
285 params: &'a Value,
286 inputs: &'a BTreeMap<String, Value>,
287 ) -> Task<'a> {
288 Task {
289 step,
290 identity,
291 params,
292 inputs,
293 }
294 }
295
296 #[tokio::test]
297 async fn add_registers_a_check_and_runs_the_operator_with_typed_params() {
298 let mut set = OperatorSet::new();
299 set.add("echo", Echo);
300 assert_eq!(set.check("echo", &table(r#"text = "hi""#)), Ok(()));
301 assert!(matches!(
302 set.check("echo", &table("")),
303 Err(OperatorError::InvalidParams(_))
304 ));
305
306 let step = Step::detached(Vec::new());
307 let identity = TaskIdentity {
308 graph: "g".into(),
309 partition: Partition::none(),
310 node: "n".into(),
311 asset: None,
312 definition: "d".into(),
313 rerun: 0,
314 };
315 let params = serde_json::json!({"text": "hi"});
316 let inputs = BTreeMap::new();
317 let task = task(&step, &identity, ¶ms, &inputs);
318 assert_eq!(
319 set.run("echo", &task, params.clone()).await.unwrap(),
320 Outcome::Succeeded(serde_json::json!({"text": "hi", "node": "n"}))
321 );
322 let err = set
323 .run("echo", &task, serde_json::json!({}))
324 .await
325 .unwrap_err();
326 assert_eq!(err.kind, taquba_workflow::StepErrorKind::Permanent);
327 let err = set.run("missing", &task, params.clone()).await.unwrap_err();
328 assert_eq!(err.kind, taquba_workflow::StepErrorKind::Permanent);
329 }
330}