Skip to main content

swale/
hook.rs

1//! The terminal hook of every pool: it writes the node's record and enqueues
2//! the event for the scheduler, and both commit with the notification's
3//! acknowledgement.
4
5use std::sync::Arc;
6
7use serde::{Deserialize, Serialize};
8use taquba::{Clock, EnqueueOptions, EnqueueRequest, MAX_KV_VALUE_SIZE};
9use taquba_workflow::{RunOutcome, StepError, TerminalEffects, TerminalHook, TerminalStatus};
10
11use crate::partition::Partition;
12use crate::records::{NodeRecord, RecordStatus};
13use crate::task::TaskIdentity;
14
15/// The queue of the scheduler's events.
16pub const EVENTS_QUEUE: &str = "swale-events";
17
18/// The event the hook enqueues for a terminated task instance.
19#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
20pub struct Event {
21    /// The graph name.
22    pub graph: String,
23    /// The partition.
24    pub partition: Partition,
25    /// The node name.
26    pub node: String,
27    /// The terminal status.
28    pub status: RecordStatus,
29}
30
31impl Event {
32    /// The JSON form of the event.
33    pub fn to_bytes(&self) -> Vec<u8> {
34        serde_json::to_vec(self).expect("an event serializes to JSON")
35    }
36
37    /// Parses the JSON form of the event.
38    pub fn from_bytes(bytes: &[u8]) -> Result<Self, serde_json::Error> {
39        serde_json::from_slice(bytes)
40    }
41}
42
43/// The terminal hook.
44#[derive(Clone)]
45pub struct RecordHook {
46    clock: Arc<dyn Clock>,
47    events_queue: String,
48}
49
50impl RecordHook {
51    /// A hook that dates records with `clock` and enqueues events on
52    /// `events_queue`.
53    pub fn new(clock: Arc<dyn Clock>, events_queue: impl Into<String>) -> Self {
54        RecordHook {
55            clock,
56            events_queue: events_queue.into(),
57        }
58    }
59
60    /// The record for `outcome`. The output is included when the record stays
61    /// within the KV value cap, and left out with `output_omitted` set
62    /// otherwise.
63    pub fn record(&self, identity: &TaskIdentity, outcome: &RunOutcome) -> NodeRecord {
64        let output = match (outcome.status, &outcome.result) {
65            (TerminalStatus::Succeeded, Some(bytes)) => {
66                Some(serde_json::from_slice(bytes).unwrap_or_else(|_| {
67                    serde_json::Value::String(String::from_utf8_lossy(bytes).into_owned())
68                }))
69            }
70            _ => None,
71        };
72        let mut record = NodeRecord {
73            status: outcome.status.into(),
74            run_id: outcome.run_id.to_string(),
75            definition: identity.definition.clone(),
76            rerun: identity.rerun,
77            terminated_at_ms: self.clock.now_ms(),
78            output,
79            output_omitted: false,
80            error: outcome.error.clone(),
81        };
82        if record.output.is_some() && record.to_bytes().len() > MAX_KV_VALUE_SIZE {
83            record.output = None;
84            record.output_omitted = true;
85        }
86        record
87    }
88}
89
90impl TerminalHook for RecordHook {
91    async fn on_termination(
92        &self,
93        outcome: &RunOutcome,
94        effects: &TerminalEffects,
95    ) -> Result<(), StepError> {
96        let identity = TaskIdentity::from_headers(&outcome.headers).map_err(|e| {
97            StepError::permanent(format!("the headers do not identify a task: {e}"))
98        })?;
99        let record = self.record(&identity, outcome);
100        effects
101            .put(identity.record_key(), record.to_bytes())
102            .map_err(|e| StepError::permanent(e.to_string()))?;
103        let event = Event {
104            graph: identity.graph,
105            partition: identity.partition,
106            node: identity.node,
107            status: record.status,
108        };
109        effects
110            .enqueue(EnqueueRequest {
111                queue: self.events_queue.clone(),
112                payload: event.to_bytes(),
113                options: EnqueueOptions::default().dedup_key(format!("evt:{}", outcome.run_id)),
114            })
115            .map_err(|e| StepError::permanent(e.to_string()))?;
116        Ok(())
117    }
118}
119
120#[cfg(test)]
121mod tests {
122    use taquba::MockClock;
123    use taquba_workflow::RunId;
124
125    use super::*;
126
127    fn identity() -> TaskIdentity {
128        TaskIdentity {
129            graph: "g".into(),
130            partition: Partition::new("20260915").unwrap(),
131            node: "extract".into(),
132            asset: Some("raw".into()),
133            definition: "d".into(),
134            rerun: 1,
135        }
136    }
137
138    fn hook() -> RecordHook {
139        RecordHook::new(Arc::new(MockClock::new(42)), EVENTS_QUEUE)
140    }
141
142    fn outcome(status: TerminalStatus, result: Option<&[u8]>, error: Option<&str>) -> RunOutcome {
143        RunOutcome {
144            run_id: RunId::new("g-20260915-extract-r1").unwrap(),
145            status,
146            result: result.map(<[u8]>::to_vec),
147            error: error.map(str::to_string),
148            headers: identity().headers(),
149            final_step: 0,
150        }
151    }
152
153    #[test]
154    fn record_copies_the_json_output_and_dates_with_the_clock() {
155        let record = hook().record(
156            &identity(),
157            &outcome(TerminalStatus::Succeeded, Some(br#"{"rows":3}"#), None),
158        );
159        assert_eq!(record.status, RecordStatus::Succeeded);
160        assert_eq!(record.run_id, "g-20260915-extract-r1");
161        assert_eq!(record.rerun, 1);
162        assert_eq!(record.terminated_at_ms, 42);
163        assert_eq!(record.output, Some(serde_json::json!({"rows": 3})));
164        assert!(!record.output_omitted);
165    }
166
167    #[test]
168    fn a_non_json_output_is_kept_as_a_string_and_a_failure_keeps_the_error() {
169        let record = hook().record(
170            &identity(),
171            &outcome(TerminalStatus::Succeeded, Some(b"plain"), None),
172        );
173        assert_eq!(
174            record.output,
175            Some(serde_json::Value::String("plain".into()))
176        );
177        let record = hook().record(
178            &identity(),
179            &outcome(TerminalStatus::Failed, None, Some("boom")),
180        );
181        assert_eq!(record.status, RecordStatus::Failed);
182        assert_eq!(record.output, None);
183        assert_eq!(record.error.as_deref(), Some("boom"));
184    }
185
186    #[test]
187    fn an_output_beyond_the_kv_cap_is_left_out_of_the_record() {
188        let big = format!("\"{}\"", "x".repeat(MAX_KV_VALUE_SIZE));
189        let record = hook().record(
190            &identity(),
191            &outcome(TerminalStatus::Succeeded, Some(big.as_bytes()), None),
192        );
193        assert_eq!(record.output, None);
194        assert!(record.output_omitted);
195        assert!(record.to_bytes().len() < MAX_KV_VALUE_SIZE);
196    }
197
198    #[test]
199    fn event_round_trips() {
200        let event = Event {
201            graph: "g".into(),
202            partition: Partition::none(),
203            node: "n".into(),
204            status: RecordStatus::Cancelled,
205        };
206        assert_eq!(Event::from_bytes(&event.to_bytes()).unwrap(), event);
207    }
208}