1use std::collections::HashMap;
16use std::fmt;
17
18use common_telemetry::{debug, error, info, warn};
19use futures::TryStreamExt;
20use serde::{Deserialize, Serialize};
21use snafu::ResultExt;
22
23use crate::error::{Result, ToJsonSnafu};
24pub(crate) use crate::store::state_store::StateStoreRef;
25use crate::{ProcedureContext, ProcedureId};
26
27pub mod poison_store;
28pub mod state_store;
29pub mod util;
30
31const PROC_PATH: &str = "procedure/";
33
34macro_rules! proc_path {
36 ($store: expr, $fmt:expr) => { format!("{}{}", $store.proc_path(), format_args!($fmt)) };
37 ($store: expr, $fmt:expr, $($args:tt)*) => { format!("{}{}", $store.proc_path(), format_args!($fmt, $($args)*)) };
38}
39
40#[cfg(test)]
41pub(crate) use proc_path;
42
43#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
45pub struct ProcedureMessage {
46 pub type_name: String,
49 pub data: String,
51 pub parent_id: Option<ProcedureId>,
53 pub step: u32,
55 #[serde(default, skip_serializing_if = "Option::is_none")]
57 pub error: Option<String>,
58 #[serde(default, skip_serializing_if = "ProcedureContext::is_empty")]
60 pub context: ProcedureContext,
61}
62
63#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
65pub struct ProcedureMessages {
66 pub messages: HashMap<ProcedureId, ProcedureMessage>,
68 pub rollback_messages: HashMap<ProcedureId, ProcedureMessage>,
70 pub finished_ids: Vec<ProcedureId>,
72}
73
74pub(crate) struct ProcedureStore {
76 proc_path: String,
77 store: StateStoreRef,
78}
79
80impl ProcedureStore {
81 pub(crate) fn new(parent_path: &str, store: StateStoreRef) -> ProcedureStore {
83 let proc_path = format!("{}{PROC_PATH}", parent_path);
84 info!("The procedure state store path is: {}", &proc_path);
85 ProcedureStore { proc_path, store }
86 }
87
88 #[inline]
89 pub(crate) fn proc_path(&self) -> &str {
90 &self.proc_path
91 }
92
93 #[cfg(test)]
95 pub(crate) async fn store_procedure(
96 &self,
97 procedure_id: ProcedureId,
98 step: u32,
99 type_name: String,
100 data: String,
101 parent_id: Option<ProcedureId>,
102 ) -> Result<()> {
103 self.store_procedure_with_context(
104 procedure_id,
105 step,
106 type_name,
107 data,
108 parent_id,
109 ProcedureContext::default(),
110 )
111 .await
112 }
113
114 pub(crate) async fn store_procedure_with_context(
115 &self,
116 procedure_id: ProcedureId,
117 step: u32,
118 type_name: String,
119 data: String,
120 parent_id: Option<ProcedureId>,
121 context: ProcedureContext,
122 ) -> Result<()> {
123 let message = ProcedureMessage {
124 type_name,
125 data,
126 parent_id,
127 step,
128 error: None,
129 context,
130 };
131 let key = ParsedKey {
132 prefix: &self.proc_path,
133 procedure_id,
134 step,
135 key_type: KeyType::Step,
136 }
137 .to_string();
138 let value = serde_json::to_string(&message).context(ToJsonSnafu)?;
139
140 self.store.put(&key, value.into_bytes()).await?;
141
142 Ok(())
143 }
144
145 pub(crate) async fn commit_procedure(
147 &self,
148 procedure_id: ProcedureId,
149 step: u32,
150 ) -> Result<()> {
151 let key = ParsedKey {
152 prefix: &self.proc_path,
153 procedure_id,
154 step,
155 key_type: KeyType::Commit,
156 }
157 .to_string();
158 self.store.put(&key, Vec::new()).await?;
159
160 Ok(())
161 }
162
163 pub(crate) async fn rollback_procedure(
165 &self,
166 procedure_id: ProcedureId,
167 message: ProcedureMessage,
168 ) -> Result<()> {
169 let key = ParsedKey {
170 prefix: &self.proc_path,
171 procedure_id,
172 step: message.step,
173 key_type: KeyType::Rollback,
174 }
175 .to_string();
176
177 self.store
178 .put(&key, serde_json::to_vec(&message).context(ToJsonSnafu)?)
179 .await?;
180
181 Ok(())
182 }
183
184 pub(crate) async fn delete_procedure(&self, procedure_id: ProcedureId) -> Result<()> {
186 let path = proc_path!(self, "{procedure_id}/");
187 let mut key_values = self.store.walk_top_down(&path).await?;
189 let mut step_keys = Vec::with_capacity(8);
191 let mut finish_keys = Vec::new();
192 while let Some((key_set, _)) = key_values.try_next().await? {
193 let key = key_set.key();
194 let Some(curr_key) = ParsedKey::parse_str(&self.proc_path, key) else {
195 warn!("Unknown key while deleting procedures, key: {}", key);
196 continue;
197 };
198 if curr_key.key_type == KeyType::Step {
199 step_keys.extend(key_set.keys());
200 } else {
201 finish_keys.extend(key_set.keys());
203 }
204 }
205
206 debug!(
207 "Delete keys for procedure {}, step_keys: {:?}, finish_keys: {:?}",
208 procedure_id, step_keys, finish_keys
209 );
210 self.store.batch_delete(step_keys.as_slice()).await?;
212 self.store.batch_delete(finish_keys.as_slice()).await?;
214 self.store.delete(&path).await?;
216 Ok(())
220 }
221
222 pub(crate) async fn load_messages(&self) -> Result<ProcedureMessages> {
224 let mut procedure_key_values: HashMap<_, (ParsedKey, Vec<u8>)> = HashMap::new();
226
227 let mut key_values = self.store.walk_top_down(&self.proc_path).await?;
229 while let Some((key_set, value)) = key_values.try_next().await? {
230 let key = key_set.key();
231 let Some(curr_key) = ParsedKey::parse_str(&self.proc_path, key) else {
232 warn!("Unknown key while loading procedures, key: {}", key);
233 continue;
234 };
235
236 if let Some(entry) = procedure_key_values.get_mut(&curr_key.procedure_id) {
237 if entry.0.step < curr_key.step {
238 entry.0 = curr_key;
239 entry.1 = value;
240 }
241 } else {
242 let _ = procedure_key_values.insert(curr_key.procedure_id, (curr_key, value));
243 }
244 }
245
246 let mut messages = HashMap::with_capacity(procedure_key_values.len());
247 let mut rollback_messages = HashMap::new();
248 let mut finished_ids = Vec::new();
249 for (procedure_id, (parsed_key, value)) in procedure_key_values {
250 match parsed_key.key_type {
251 KeyType::Step => {
252 let Some(message) = self.load_one_message(&parsed_key, &value) else {
253 continue;
256 };
257 let _ = messages.insert(procedure_id, message);
258 }
259 KeyType::Commit => {
260 finished_ids.push(procedure_id);
261 }
262 KeyType::Rollback => {
263 let Some(message) = self.load_one_message(&parsed_key, &value) else {
264 continue;
267 };
268 let _ = rollback_messages.insert(procedure_id, message);
269 }
270 }
271 }
272
273 Ok(ProcedureMessages {
274 messages,
275 rollback_messages,
276 finished_ids,
277 })
278 }
279
280 fn load_one_message(&self, key: &ParsedKey, value: &[u8]) -> Option<ProcedureMessage> {
281 serde_json::from_slice(value)
282 .map_err(|e| {
283 error!("Failed to parse value, key: {:?}, source: {:?}", key, e);
285 e
286 })
287 .ok()
288 }
289}
290
291#[derive(Debug, PartialEq, Eq)]
293enum KeyType {
294 Step,
295 Commit,
296 Rollback,
297}
298
299impl KeyType {
300 fn as_str(&self) -> &'static str {
301 match self {
302 KeyType::Step => "step",
303 KeyType::Commit => "commit",
304 KeyType::Rollback => "rollback",
305 }
306 }
307
308 fn from_str(s: &str) -> Option<KeyType> {
309 match s {
310 "step" => Some(KeyType::Step),
311 "commit" => Some(KeyType::Commit),
312 "rollback" => Some(KeyType::Rollback),
313 _ => None,
314 }
315 }
316}
317
318#[derive(Debug, PartialEq, Eq)]
320struct ParsedKey<'a> {
321 prefix: &'a str,
322 procedure_id: ProcedureId,
323 step: u32,
324 key_type: KeyType,
325}
326
327impl fmt::Display for ParsedKey<'_> {
328 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
329 write!(
330 f,
331 "{}{}/{:010}.{}",
332 self.prefix,
333 self.procedure_id,
334 self.step,
335 self.key_type.as_str(),
336 )
337 }
338}
339
340impl<'a> ParsedKey<'a> {
341 fn parse_str(prefix: &'a str, input: &str) -> Option<ParsedKey<'a>> {
343 let input = input.strip_prefix(prefix)?;
344 let mut iter = input.rsplit('/');
345 let name = iter.next()?;
346 let id_str = iter.next()?;
347
348 let procedure_id = ProcedureId::parse_str(id_str).ok()?;
349
350 let mut parts = name.split('.');
351 let step_str = parts.next()?;
352 let suffix = parts.next()?;
353 let key_type = KeyType::from_str(suffix)?;
354 let step = step_str.parse().ok()?;
355
356 Some(ParsedKey {
357 prefix,
358 procedure_id,
359 step,
360 key_type,
361 })
362 }
363}
364
365#[cfg(test)]
366mod tests {
367 use std::sync::Arc;
368
369 use object_store::ObjectStore;
370
371 use crate::BoxedProcedure;
372 use crate::procedure::PoisonKeys;
373 use crate::store::state_store::ObjectStateStore;
374
375 impl ProcedureStore {
376 pub(crate) fn from_object_store(store: ObjectStore) -> ProcedureStore {
377 let state_store = ObjectStateStore::new(store);
378
379 ProcedureStore::new("data/", Arc::new(state_store))
380 }
381 }
382
383 use async_trait::async_trait;
384 use common_event_recorder::{PersistentEventContext, TriggerReason};
385 use common_test_util::temp_dir::{TempDir, create_temp_dir};
386 use object_store::services::Fs as Builder;
387
388 use super::*;
389 use crate::{Context, LockKey, Procedure, Status};
390
391 fn procedure_store_for_test(dir: &TempDir) -> ProcedureStore {
392 let store_dir = dir.path().to_str().unwrap();
393 let builder = Builder::default().root(store_dir);
394 let object_store = ObjectStore::new(builder).unwrap();
395
396 ProcedureStore::from_object_store(object_store)
397 }
398
399 #[test]
400 fn test_parsed_key() {
401 let dir = create_temp_dir("store_procedure");
402 let store = procedure_store_for_test(&dir);
403
404 let procedure_id = ProcedureId::random();
405 let key = ParsedKey {
406 prefix: &store.proc_path,
407 procedure_id,
408 step: 2,
409 key_type: KeyType::Step,
410 };
411 assert_eq!(
412 proc_path!(store, "{procedure_id}/0000000002.step"),
413 key.to_string()
414 );
415 assert_eq!(
416 key,
417 ParsedKey::parse_str(&store.proc_path, &key.to_string()).unwrap()
418 );
419
420 let key = ParsedKey {
421 prefix: &store.proc_path,
422 procedure_id,
423 step: 2,
424 key_type: KeyType::Commit,
425 };
426 assert_eq!(
427 proc_path!(store, "{procedure_id}/0000000002.commit"),
428 key.to_string()
429 );
430 assert_eq!(
431 key,
432 ParsedKey::parse_str(&store.proc_path, &key.to_string()).unwrap()
433 );
434
435 let key = ParsedKey {
436 prefix: &store.proc_path,
437 procedure_id,
438 step: 2,
439 key_type: KeyType::Rollback,
440 };
441 assert_eq!(
442 proc_path!(store, "{procedure_id}/0000000002.rollback"),
443 key.to_string()
444 );
445 assert_eq!(
446 key,
447 ParsedKey::parse_str(&store.proc_path, &key.to_string()).unwrap()
448 );
449 }
450
451 #[test]
452 fn test_parse_invalid_key() {
453 let dir = create_temp_dir("store_procedure");
454 let store = procedure_store_for_test(&dir);
455
456 assert!(ParsedKey::parse_str(&store.proc_path, "").is_none());
457 assert!(ParsedKey::parse_str(&store.proc_path, "invalidprefix").is_none());
458 assert!(ParsedKey::parse_str(&store.proc_path, "procedu/0000000003.step").is_none());
459 assert!(ParsedKey::parse_str(&store.proc_path, "procedure-0000000003.step").is_none());
460
461 let procedure_id = ProcedureId::random();
462 let input = proc_path!(store, "{procedure_id}");
463 assert!(ParsedKey::parse_str(&store.proc_path, &input).is_none());
464
465 let input = proc_path!(store, "{procedure_id}");
466 assert!(ParsedKey::parse_str(&store.proc_path, &input).is_none());
467
468 let input = proc_path!(store, "{procedure_id}/0000000003");
469 assert!(ParsedKey::parse_str(&store.proc_path, &input).is_none());
470
471 let input = proc_path!(store, "{procedure_id}/0000000003.");
472 assert!(ParsedKey::parse_str(&store.proc_path, &input).is_none());
473
474 let input = proc_path!(store, "{procedure_id}/0000000003.other");
475 assert!(ParsedKey::parse_str(&store.proc_path, &input).is_none());
476
477 assert!(ParsedKey::parse_str(&store.proc_path, "12345/0000000003.step").is_none());
478
479 let input = proc_path!(store, "{procedure_id}-0000000003.commit");
480 assert!(ParsedKey::parse_str(&store.proc_path, &input).is_none());
481 }
482
483 #[test]
484 fn test_procedure_message() {
485 let mut message = ProcedureMessage {
486 type_name: "TestMessage".to_string(),
487 data: "no parent id".to_string(),
488 parent_id: None,
489 step: 4,
490 error: None,
491 context: ProcedureContext::default(),
492 };
493
494 let json = serde_json::to_string(&message).unwrap();
495 assert_eq!(
496 json,
497 r#"{"type_name":"TestMessage","data":"no parent id","parent_id":null,"step":4}"#
498 );
499
500 let procedure_id = ProcedureId::parse_str("9f805a1f-05f7-490c-9f91-bd56e3cc54c1").unwrap();
501 message.parent_id = Some(procedure_id);
502 let json = serde_json::to_string(&message).unwrap();
503 assert_eq!(
504 json,
505 r#"{"type_name":"TestMessage","data":"no parent id","parent_id":"9f805a1f-05f7-490c-9f91-bd56e3cc54c1","step":4}"#
506 );
507
508 message.context = ProcedureContext::from_event_context(PersistentEventContext::new(
509 TriggerReason::AutoRebalance,
510 ));
511 let json = serde_json::to_string(&message).unwrap();
512 assert_eq!(
513 json,
514 r#"{"type_name":"TestMessage","data":"no parent id","parent_id":"9f805a1f-05f7-490c-9f91-bd56e3cc54c1","step":4,"context":{"event_context":{"reason":"auto_rebalance"}}}"#
515 );
516
517 message.context.actor = Some("alice".to_string());
518 let json = serde_json::to_string(&message).unwrap();
519 assert!(json.contains(r#""actor":"alice""#));
520 assert_eq!(
521 serde_json::from_str::<ProcedureMessage>(&json).unwrap(),
522 message
523 );
524
525 let legacy: ProcedureMessage = serde_json::from_str(
526 r#"{"type_name":"TestMessage","data":"legacy","parent_id":null,"step":1}"#,
527 )
528 .unwrap();
529 assert_eq!(legacy.context, ProcedureContext::default());
530 }
531
532 struct MockProcedure {
533 data: String,
534 }
535
536 impl MockProcedure {
537 fn new(data: impl Into<String>) -> MockProcedure {
538 MockProcedure { data: data.into() }
539 }
540 }
541
542 #[async_trait]
543 impl Procedure for MockProcedure {
544 fn type_name(&self) -> &str {
545 "MockProcedure"
546 }
547
548 async fn execute(&mut self, _ctx: &Context) -> Result<Status> {
549 unimplemented!()
550 }
551
552 fn dump(&self) -> Result<String> {
553 Ok(self.data.clone())
554 }
555
556 fn lock_key(&self) -> LockKey {
557 LockKey::default()
558 }
559
560 fn poison_keys(&self) -> PoisonKeys {
561 PoisonKeys::default()
562 }
563 }
564
565 #[tokio::test]
566 async fn test_store_procedure() {
567 let dir = create_temp_dir("store_procedure");
568 let store = procedure_store_for_test(&dir);
569
570 let procedure_id = ProcedureId::random();
571 let procedure: BoxedProcedure = Box::new(MockProcedure::new("test store procedure"));
572 let type_name = procedure.type_name().to_string();
573 let data = procedure.dump().unwrap();
574 store
575 .store_procedure(procedure_id, 0, type_name, data, None)
576 .await
577 .unwrap();
578
579 let ProcedureMessages {
580 messages,
581 rollback_messages,
582 finished_ids,
583 } = store.load_messages().await.unwrap();
584 assert_eq!(1, messages.len());
585 assert!(rollback_messages.is_empty());
586 assert!(finished_ids.is_empty());
587 let msg = messages.get(&procedure_id).unwrap();
588 let expect = ProcedureMessage {
589 type_name: "MockProcedure".to_string(),
590 data: "test store procedure".to_string(),
591 parent_id: None,
592 step: 0,
593 error: None,
594 context: ProcedureContext::default(),
595 };
596 assert_eq!(expect, *msg);
597 }
598
599 #[tokio::test]
600 async fn test_commit_procedure() {
601 let dir = create_temp_dir("commit_procedure");
602 let store = procedure_store_for_test(&dir);
603
604 let procedure_id = ProcedureId::random();
605 let procedure: BoxedProcedure = Box::new(MockProcedure::new("test store procedure"));
606 let type_name = procedure.type_name().to_string();
607 let data = procedure.dump().unwrap();
608 store
609 .store_procedure(procedure_id, 0, type_name, data, None)
610 .await
611 .unwrap();
612 store.commit_procedure(procedure_id, 1).await.unwrap();
613
614 let ProcedureMessages {
615 messages,
616 rollback_messages,
617 finished_ids,
618 } = store.load_messages().await.unwrap();
619 assert!(messages.is_empty());
620 assert!(rollback_messages.is_empty());
621 assert_eq!(&[procedure_id], &finished_ids[..]);
622 }
623
624 #[tokio::test]
625 async fn test_rollback_procedure() {
626 let dir = create_temp_dir("rollback_procedure");
627 let store = procedure_store_for_test(&dir);
628
629 let procedure_id = ProcedureId::random();
630 let procedure: BoxedProcedure = Box::new(MockProcedure::new("test store procedure"));
631 let type_name = procedure.type_name().to_string();
632 let data = procedure.dump().unwrap();
633 store
634 .store_procedure(procedure_id, 0, type_name.clone(), data.clone(), None)
635 .await
636 .unwrap();
637 let message = ProcedureMessage {
638 type_name,
639 data,
640 parent_id: None,
641 step: 1,
642 error: None,
643 context: ProcedureContext::default(),
644 };
645 store
646 .rollback_procedure(procedure_id, message)
647 .await
648 .unwrap();
649
650 let ProcedureMessages {
651 messages,
652 rollback_messages,
653 finished_ids,
654 } = store.load_messages().await.unwrap();
655 assert!(messages.is_empty());
656 assert_eq!(1, rollback_messages.len());
657 assert!(finished_ids.is_empty());
658 assert!(rollback_messages.contains_key(&procedure_id));
659 }
660
661 #[tokio::test]
662 async fn test_delete_procedure() {
663 let dir = create_temp_dir("delete_procedure");
664 let store = procedure_store_for_test(&dir);
665
666 let procedure_id = ProcedureId::random();
667 let procedure: BoxedProcedure = Box::new(MockProcedure::new("test store procedure"));
668 let type_name = procedure.type_name().to_string();
669 let data = procedure.dump().unwrap();
670 store
671 .store_procedure(procedure_id, 0, type_name, data, None)
672 .await
673 .unwrap();
674 let type_name = procedure.type_name().to_string();
675 let data = procedure.dump().unwrap();
676 store
677 .store_procedure(procedure_id, 1, type_name, data, None)
678 .await
679 .unwrap();
680
681 store.delete_procedure(procedure_id).await.unwrap();
682
683 let ProcedureMessages {
684 messages,
685 rollback_messages,
686 finished_ids,
687 } = store.load_messages().await.unwrap();
688 assert!(messages.is_empty());
689 assert!(rollback_messages.is_empty());
690 assert!(finished_ids.is_empty());
691 }
692
693 #[tokio::test]
694 async fn test_delete_committed_procedure() {
695 let dir = create_temp_dir("delete_committed");
696 let store = procedure_store_for_test(&dir);
697
698 let procedure_id = ProcedureId::random();
699 let procedure: BoxedProcedure = Box::new(MockProcedure::new("test store procedure"));
700
701 let type_name = procedure.type_name().to_string();
702 let data = procedure.dump().unwrap();
703 store
704 .store_procedure(procedure_id, 0, type_name, data, None)
705 .await
706 .unwrap();
707
708 let type_name = procedure.type_name().to_string();
709 let data = procedure.dump().unwrap();
710 store
711 .store_procedure(procedure_id, 1, type_name, data, None)
712 .await
713 .unwrap();
714 store.commit_procedure(procedure_id, 2).await.unwrap();
715
716 store.delete_procedure(procedure_id).await.unwrap();
717
718 let ProcedureMessages {
719 messages,
720 rollback_messages,
721 finished_ids,
722 } = store.load_messages().await.unwrap();
723 assert!(messages.is_empty());
724 assert!(rollback_messages.is_empty());
725 assert!(finished_ids.is_empty());
726 }
727
728 #[tokio::test]
729 async fn test_load_messages() {
730 let dir = create_temp_dir("load_messages");
731 let store = procedure_store_for_test(&dir);
732
733 let id0 = ProcedureId::random();
735 let procedure: BoxedProcedure = Box::new(MockProcedure::new("id0-0"));
736 let type_name = procedure.type_name().to_string();
737 let data = procedure.dump().unwrap();
738 store
739 .store_procedure(id0, 0, type_name, data, None)
740 .await
741 .unwrap();
742 let procedure: BoxedProcedure = Box::new(MockProcedure::new("id0-1"));
743 let type_name = procedure.type_name().to_string();
744 let data = procedure.dump().unwrap();
745 store
746 .store_procedure(id0, 1, type_name, data, None)
747 .await
748 .unwrap();
749 let procedure: BoxedProcedure = Box::new(MockProcedure::new("id0-2"));
750 let type_name = procedure.type_name().to_string();
751 let data = procedure.dump().unwrap();
752 store
753 .store_procedure(id0, 2, type_name, data, None)
754 .await
755 .unwrap();
756
757 let id1 = ProcedureId::random();
759 let procedure: BoxedProcedure = Box::new(MockProcedure::new("id1-0"));
760 let type_name = procedure.type_name().to_string();
761 let data = procedure.dump().unwrap();
762 store
763 .store_procedure(id1, 0, type_name, data, None)
764 .await
765 .unwrap();
766 let procedure: BoxedProcedure = Box::new(MockProcedure::new("id1-1"));
767 let type_name = procedure.type_name().to_string();
768 let data = procedure.dump().unwrap();
769 store
770 .store_procedure(id1, 1, type_name, data, None)
771 .await
772 .unwrap();
773 store.commit_procedure(id1, 2).await.unwrap();
774
775 let id2 = ProcedureId::random();
777 let procedure: BoxedProcedure = Box::new(MockProcedure::new("id2-0"));
778 let type_name = procedure.type_name().to_string();
779 let data = procedure.dump().unwrap();
780 store
781 .store_procedure(id2, 0, type_name, data, None)
782 .await
783 .unwrap();
784
785 let ProcedureMessages {
786 messages,
787 rollback_messages,
788 finished_ids,
789 } = store.load_messages().await.unwrap();
790 assert_eq!(2, messages.len());
791 assert!(rollback_messages.is_empty());
792 assert_eq!(1, finished_ids.len());
793
794 let msg = messages.get(&id0).unwrap();
795 assert_eq!("id0-2", msg.data);
796 let msg = messages.get(&id2).unwrap();
797 assert_eq!("id2-0", msg.data);
798 }
799}