Skip to main content

stasis/application/runtime/
coordinator_failover_job_handler.rs

1use std::collections::BTreeSet;
2use std::sync::Arc;
3
4use async_trait::async_trait;
5use chrono::Utc;
6use serde::Deserialize;
7use serde_json::json;
8
9use crate::application::runtime::in_memory_runtime::{JobExecutionOutcome, JobHandler};
10use crate::domain::errors::Result;
11use crate::domain::runtime::cluster_node::ClusterNodeHeartbeat;
12use crate::domain::runtime::job::Job;
13use crate::ports::outbound::runtime::cluster_node_store::ClusterNodeStore;
14
15#[derive(Clone)]
16pub struct CoordinatorFailoverJobHandler {
17    cluster_store: Arc<dyn ClusterNodeStore>,
18}
19
20#[derive(Deserialize)]
21struct CoordinatorFailoverPayload {
22    coordinator_node_id: String,
23    failover_to_node_id: Option<String>,
24    queue_scope: Option<Vec<String>>,
25    reason: Option<String>,
26}
27
28impl CoordinatorFailoverJobHandler {
29    pub fn new(cluster_store: Arc<dyn ClusterNodeStore>) -> Self {
30        Self { cluster_store }
31    }
32
33    fn parse_payload(raw: &str) -> std::result::Result<CoordinatorFailoverPayload, String> {
34        let payload: CoordinatorFailoverPayload = serde_json::from_str(raw)
35            .map_err(|err| format!("invalid coordinator failover payload json: {err}"))?;
36
37        if payload.coordinator_node_id.trim().is_empty() {
38            return Err("coordinator_node_id must not be empty".to_string());
39        }
40
41        if let Some(target) = &payload.failover_to_node_id
42            && target.trim().is_empty()
43        {
44            return Err("failover_to_node_id must not be empty when provided".to_string());
45        }
46
47        Ok(payload)
48    }
49
50    fn failure(message: impl Into<String>) -> JobExecutionOutcome {
51        JobExecutionOutcome::FatalFailure {
52            message: message.into(),
53            execution_id: None,
54            diagnostics: None,
55        }
56    }
57
58    fn retryable(message: impl Into<String>) -> JobExecutionOutcome {
59        JobExecutionOutcome::RetryableFailure {
60            message: message.into(),
61            execution_id: None,
62            diagnostics: None,
63        }
64    }
65
66    fn remaining_ttl_seconds(
67        lease_expires_at: chrono::DateTime<Utc>,
68        now: chrono::DateTime<Utc>,
69    ) -> i64 {
70        lease_expires_at
71            .signed_duration_since(now)
72            .num_seconds()
73            .max(1)
74    }
75}
76
77#[async_trait]
78impl JobHandler for CoordinatorFailoverJobHandler {
79    fn job_type(&self) -> &'static str {
80        "workflow.stasis.cluster.coordinator_failover"
81    }
82
83    async fn execute(&self, job: &Job) -> Result<JobExecutionOutcome> {
84        let payload = match Self::parse_payload(&job.payload_ref) {
85            Ok(payload) => payload,
86            Err(message) => return Ok(Self::failure(message)),
87        };
88
89        let now = Utc::now();
90        let nodes = match self.cluster_store.list().await {
91            Ok(nodes) => nodes,
92            Err(err) => {
93                return Ok(Self::retryable(format!(
94                    "failed to load cluster nodes for failover: {err}"
95                )));
96            }
97        };
98
99        let Some(source) = nodes
100            .iter()
101            .find(|node| node.node_id == payload.coordinator_node_id)
102            .cloned()
103        else {
104            return Ok(Self::failure(format!(
105                "coordinator node not found: {}",
106                payload.coordinator_node_id
107            )));
108        };
109
110        if source.lease_expires_at < now {
111            return Ok(Self::failure(format!(
112                "coordinator node is not active: {}",
113                source.node_id
114            )));
115        }
116
117        let target_node_id = payload.failover_to_node_id.clone().or_else(|| {
118            nodes
119                .iter()
120                .filter(|node| node.node_id != source.node_id)
121                .filter(|node| node.lease_expires_at >= now)
122                .find(|node| node.region == source.region)
123                .map(|node| node.node_id.clone())
124        });
125
126        let Some(target_node_id) = target_node_id else {
127            return Ok(Self::failure("no active failover target available"));
128        };
129
130        if target_node_id == source.node_id {
131            return Ok(Self::failure(
132                "failover target must differ from coordinator node",
133            ));
134        }
135
136        let Some(target) = nodes
137            .iter()
138            .find(|node| node.node_id == target_node_id)
139            .cloned()
140        else {
141            return Ok(Self::failure(format!(
142                "failover target node not found: {}",
143                target_node_id
144            )));
145        };
146
147        if target.lease_expires_at < now {
148            return Ok(Self::failure(format!(
149                "failover target node is not active: {}",
150                target.node_id
151            )));
152        }
153
154        let moved_queues = if let Some(scope) = payload.queue_scope {
155            scope.into_iter().collect::<BTreeSet<_>>()
156        } else {
157            source
158                .queue_ownership
159                .iter()
160                .cloned()
161                .collect::<BTreeSet<_>>()
162        };
163
164        let source_queues = source
165            .queue_ownership
166            .into_iter()
167            .filter(|queue| !moved_queues.contains(queue))
168            .collect::<Vec<_>>();
169
170        let mut target_queue_set = target.queue_ownership.into_iter().collect::<BTreeSet<_>>();
171        for queue in &moved_queues {
172            target_queue_set.insert(queue.clone());
173        }
174
175        let source_ttl = Self::remaining_ttl_seconds(source.lease_expires_at, now);
176        let target_ttl = Self::remaining_ttl_seconds(target.lease_expires_at, now);
177
178        let source_update = self
179            .cluster_store
180            .heartbeat(ClusterNodeHeartbeat {
181                node_id: source.node_id.clone(),
182                heartbeat_at: now,
183                lease_ttl_seconds: source_ttl,
184                queue_ownership: Some(source_queues),
185                capability_tags: None,
186                metadata: Some(
187                    json!({
188                        "action": "coordinator_failover_source",
189                        "reason": payload.reason,
190                    })
191                    .to_string(),
192                ),
193            })
194            .await;
195
196        if source_update.is_err() {
197            return Ok(Self::retryable(
198                "failed to update source node during failover",
199            ));
200        }
201
202        let target_update = self
203            .cluster_store
204            .heartbeat(ClusterNodeHeartbeat {
205                node_id: target.node_id.clone(),
206                heartbeat_at: now,
207                lease_ttl_seconds: target_ttl,
208                queue_ownership: Some(target_queue_set.into_iter().collect::<Vec<_>>()),
209                capability_tags: None,
210                metadata: Some(
211                    json!({
212                        "action": "coordinator_failover_target",
213                        "from_node": source.node_id,
214                        "reason": payload.reason,
215                    })
216                    .to_string(),
217                ),
218            })
219            .await;
220
221        if target_update.is_err() {
222            return Ok(Self::retryable(
223                "failed to update target node during failover",
224            ));
225        }
226
227        let diagnostics = json!({
228            "status": "success",
229            "source_node": payload.coordinator_node_id,
230            "target_node": target.node_id,
231            "moved_queues": moved_queues.into_iter().collect::<Vec<_>>(),
232        })
233        .to_string();
234
235        Ok(JobExecutionOutcome::Success {
236            sttp_output_node_id: "sttp:out:stasis:cluster:coordinator_failover".to_string(),
237            execution_id: None,
238            diagnostics: Some(diagnostics),
239        })
240    }
241}
242
243#[cfg(test)]
244mod tests {
245    use chrono::{Duration, Utc};
246
247    use crate::application::runtime::coordinator_failover_job_handler::CoordinatorFailoverJobHandler;
248    use crate::application::runtime::in_memory_runtime::JobHandler;
249    use crate::domain::runtime::cluster_node::{ClusterNodeRole, NewClusterNode};
250    use crate::domain::runtime::job::{BackoffPolicy, Job, JobState};
251    use crate::infrastructure::runtime::in_memory_cluster_node_store::InMemoryClusterNodeStore;
252    use crate::ports::outbound::runtime::cluster_node_store::ClusterNodeStore;
253    use std::sync::Arc;
254
255    fn sample_job(payload_ref: String) -> Job {
256        Job {
257            id: "job.cluster.failover.1".to_string(),
258            queue: "cluster-control".to_string(),
259            job_type: "workflow.stasis.cluster.coordinator_failover".to_string(),
260            payload_ref,
261            state: JobState::Enqueued,
262            priority: 100,
263            attempts: 0,
264            max_attempts: 3,
265            backoff_policy: BackoffPolicy::default(),
266            idempotency_key: "idem-failover".to_string(),
267            correlation_id: "corr-failover".to_string(),
268            causation_id: "cause-failover".to_string(),
269            trace_id: "trace-failover".to_string(),
270            sttp_input_node_id: "sttp:in:cluster:failover".to_string(),
271            sttp_output_node_id: None,
272            lease_owner: None,
273            lease_expires_at: None,
274            heartbeat_at: None,
275            scheduled_at: Utc::now(),
276            started_at: None,
277            finished_at: None,
278            last_error: None,
279        }
280    }
281
282    #[tokio::test]
283    async fn handler_moves_queues_to_target_node() {
284        let store = InMemoryClusterNodeStore::default();
285        let now = Utc::now();
286
287        store
288            .register(NewClusterNode {
289                node_id: "node.coord.a".to_string(),
290                role: ClusterNodeRole::Coordinator,
291                region: "us-east".to_string(),
292                queue_ownership: vec!["default".to_string(), "priority".to_string()],
293                capability_tags: vec![],
294                heartbeat_at: now,
295                lease_ttl_seconds: 60,
296                metadata: None,
297            })
298            .await
299            .expect("register source should succeed");
300
301        store
302            .register(NewClusterNode {
303                node_id: "node.coord.b".to_string(),
304                role: ClusterNodeRole::Coordinator,
305                region: "us-east".to_string(),
306                queue_ownership: vec![],
307                capability_tags: vec![],
308                heartbeat_at: now,
309                lease_ttl_seconds: 60,
310                metadata: None,
311            })
312            .await
313            .expect("register target should succeed");
314
315        let handler = CoordinatorFailoverJobHandler::new(Arc::new(store.clone()));
316        let payload = serde_json::json!({
317            "coordinator_node_id": "node.coord.a",
318            "failover_to_node_id": "node.coord.b",
319            "queue_scope": ["priority"],
320            "reason": "planned",
321        })
322        .to_string();
323
324        let outcome = handler
325            .execute(&sample_job(payload))
326            .await
327            .expect("handler execution should succeed");
328
329        assert!(matches!(
330            outcome,
331            crate::application::runtime::in_memory_runtime::JobExecutionOutcome::Success { .. }
332        ));
333
334        let source = store
335            .get("node.coord.a")
336            .await
337            .expect("source get should succeed")
338            .expect("source should exist");
339        let target = store
340            .get("node.coord.b")
341            .await
342            .expect("target get should succeed")
343            .expect("target should exist");
344
345        assert_eq!(source.queue_ownership, vec!["default".to_string()]);
346        assert!(target.queue_ownership.iter().any(|q| q == "priority"));
347    }
348
349    #[tokio::test]
350    async fn handler_returns_fatal_failure_for_missing_target() {
351        let store = InMemoryClusterNodeStore::default();
352        let now = Utc::now();
353
354        store
355            .register(NewClusterNode {
356                node_id: "node.coord.a".to_string(),
357                role: ClusterNodeRole::Coordinator,
358                region: "us-east".to_string(),
359                queue_ownership: vec!["default".to_string()],
360                capability_tags: vec![],
361                heartbeat_at: now - Duration::seconds(1),
362                lease_ttl_seconds: 60,
363                metadata: None,
364            })
365            .await
366            .expect("register source should succeed");
367
368        let handler = CoordinatorFailoverJobHandler::new(Arc::new(store));
369        let payload = serde_json::json!({
370            "coordinator_node_id": "node.coord.a",
371            "failover_to_node_id": "missing",
372        })
373        .to_string();
374
375        let outcome = handler
376            .execute(&sample_job(payload))
377            .await
378            .expect("handler execution should succeed");
379
380        assert!(matches!(
381            outcome,
382            crate::application::runtime::in_memory_runtime::JobExecutionOutcome::FatalFailure { .. }
383        ));
384    }
385}