Skip to main content

swale/
operator.rs

1//! The operators a definition can name: the check of their parameters at load
2//! time and their execution at run time.
3//!
4//! An [`OperatorSet`] maps an operator name to a parameter check and, for an
5//! operator the process runs, to an [`Operator`]. [`OperatorSet::builtin`]
6//! contains the operators of this crate, and a consumer adds its own with
7//! [`OperatorSet::add`] or, for a check alone, [`OperatorSet::register`].
8
9use 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/// One task instance as an operator sees it.
24#[derive(Debug, Clone, Copy)]
25pub struct Task<'a> {
26    /// The step of the run, with the lease, the cancellation token and the
27    /// attempt count.
28    pub step: &'a Step,
29    /// The identity of the task instance.
30    pub identity: &'a TaskIdentity,
31    /// The node's parameters with every template rendered.
32    pub params: &'a Value,
33    /// The output of each upstream node with a succeeded record.
34    pub inputs: &'a BTreeMap<String, Value>,
35}
36
37/// The outcome of an operator. An infrastructure failure is a [`StepError`]
38/// instead: transient for a retry, permanent for a dead-letter.
39#[derive(Debug, Clone, PartialEq, Eq)]
40pub enum Outcome {
41    /// The task instance succeeded with this output.
42    Succeeded(Value),
43    /// The task instance failed with this reason.
44    Failed(String),
45}
46
47/// An operator the process runs.
48pub trait Operator: Send + Sync + 'static {
49    /// The parameter type. A node's `params` table must deserialize into it
50    /// at load time, and the rendered parameters do so at run time.
51    type Params: DeserializeOwned + Send;
52
53    /// Runs one task instance. The implementation must be idempotent per
54    /// attempt, because delivery is at least once.
55    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/// The registered operators.
72#[derive(Default)]
73pub struct OperatorSet {
74    entries: BTreeMap<String, Entry>,
75}
76
77/// The reason a node's operator or parameters are rejected.
78#[derive(Debug, Clone, PartialEq, Eq)]
79pub enum OperatorError {
80    /// The operator is not registered.
81    Unknown,
82    /// The parameters do not deserialize into the operator's type. The
83    /// string is the deserializer's message.
84    InvalidParams(String),
85}
86
87impl OperatorSet {
88    /// An empty set.
89    pub fn new() -> Self {
90        Self::default()
91    }
92
93    /// The operators of this crate: `subprocess`.
94    pub fn builtin() -> Self {
95        let mut set = Self::new();
96        set.add("subprocess", Subprocess::default());
97        set
98    }
99
100    /// Registers `name` with `P` as its parameter type and without an
101    /// operator to run. A second registration of `name` replaces the first.
102    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    /// Adds `operator` as `name`, with its parameter type as the check. A
113    /// second entry for `name` replaces the first.
114    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    /// Whether `name` is registered.
135    pub fn contains(&self, name: &str) -> bool {
136        self.entries.contains_key(name)
137    }
138
139    /// Whether `name` has an operator to run.
140    pub fn runs(&self, name: &str) -> bool {
141        self.entries.get(name).is_some_and(|e| e.run.is_some())
142    }
143
144    /// The registered names in lexical order.
145    pub fn names(&self) -> impl Iterator<Item = &str> {
146        self.entries.keys().map(String::as_str)
147    }
148
149    /// Checks a node's parameters against the operator `name`.
150    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    /// Runs `task` with the operator `name` and the rendered `params`. An
156    /// operator that is unknown or checked only is a permanent error.
157    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/// Parameters of the `subprocess` operator: a program and its arguments.
189#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
190#[serde(deny_unknown_fields)]
191pub struct SubprocessParams {
192    /// The program followed by its arguments. It must not be empty.
193    #[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, &params, &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}