Skip to main content

sz_rust_workflow/engine/
state_machine.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2024-2026 SZ-Rust Team
3//
4use std::sync::Arc;
5
6use chrono::Utc;
7
8use crate::definition::StateMachineDefinition;
9use crate::error::{WorkflowError, WorkflowErrorCode, WorkflowResult};
10use crate::guard::GuardEvaluator;
11use crate::instance::InstanceStatus;
12use crate::observability::{AuditLogger, WorkflowEvent, WorkflowEventBus};
13use crate::repository::InstanceRepository;
14
15/// 迁移结果。
16#[derive(Debug, Clone)]
17pub struct TransitionResult {
18    pub from_state: String,
19    pub to_state: String,
20    pub event: String,
21    pub migrated: bool,
22}
23
24/// 状态机引擎。
25pub struct StateMachineEngine {
26    instance_repo: Arc<dyn InstanceRepository>,
27    guard_evaluator: Arc<dyn GuardEvaluator>,
28    event_bus: Arc<dyn WorkflowEventBus>,
29    audit: Arc<AuditLogger>,
30}
31
32impl StateMachineEngine {
33    pub fn new(
34        instance_repo: Arc<dyn InstanceRepository>,
35        guard_evaluator: Arc<dyn GuardEvaluator>,
36        event_bus: Arc<dyn WorkflowEventBus>,
37        audit: Arc<AuditLogger>,
38    ) -> Self {
39        Self {
40            instance_repo,
41            guard_evaluator,
42            event_bus,
43            audit,
44        }
45    }
46
47    /// 触发事件,执行状态迁移。
48    pub async fn fire(
49        &self,
50        instance_id: &str,
51        event: &str,
52        machine: &StateMachineDefinition,
53        payload: serde_json::Value,
54    ) -> WorkflowResult<TransitionResult> {
55        let instance = self.instance_repo.get(instance_id).await?.ok_or_else(|| {
56            WorkflowError::with_field(
57                WorkflowErrorCode::InstanceNotFound,
58                "实例不存在",
59                "instance_id",
60                instance_id,
61            )
62        })?;
63
64        if !instance.status.is_handleable() {
65            return Err(WorkflowError::with_field(
66                WorkflowErrorCode::InstanceNotHandleable,
67                format!("实例状态 {} 不可办理", instance.status),
68                "status",
69                &instance.status.to_string(),
70            ));
71        }
72
73        let current_state = instance
74            .context
75            .get("current_state")
76            .and_then(|v| v.as_str())
77            .unwrap_or(&machine.initial_state)
78            .to_string();
79
80        let transition = machine
81            .transitions
82            .iter()
83            .find(|t| t.from == current_state && t.event == event)
84            .ok_or_else(|| {
85                WorkflowError::with_field(
86                    WorkflowErrorCode::NoMatchingTransition,
87                    format!("状态 {} 不接受事件 {}", current_state, event),
88                    "state",
89                    &current_state,
90                )
91            })?;
92
93        if let Some(ref guard) = transition.guard {
94            let mut ctx = instance.context.clone();
95            if let serde_json::Value::Object(ref mut obj) = ctx {
96                if let serde_json::Value::Object(ref payload_obj) = payload {
97                    for (k, v) in payload_obj {
98                        obj.insert(k.clone(), v.clone());
99                    }
100                }
101                obj.insert("payload".into(), payload.clone());
102            }
103            let guard_result = self.guard_evaluator.evaluate(guard, &ctx).await;
104            match guard_result {
105                Ok(true) => {}
106                Ok(false) => {
107                    return Ok(TransitionResult {
108                        from_state: current_state.clone(),
109                        to_state: current_state,
110                        event: event.into(),
111                        migrated: false,
112                    });
113                }
114                Err(e) => return Err(e),
115            }
116        }
117
118        let expected_version = instance.version_lock;
119        let mut updated = instance.clone();
120        if let serde_json::Value::Object(ref mut obj) = updated.context {
121            obj.insert(
122                "current_state".into(),
123                serde_json::Value::String(transition.to.clone()),
124            );
125            obj.insert("payload".into(), payload);
126        }
127        updated.bump_version();
128
129        let success = self
130            .instance_repo
131            .update_with_version(&updated, expected_version)
132            .await?;
133        if !success {
134            return Err(WorkflowError::new(
135                WorkflowErrorCode::OptimisticLockConflict,
136                "乐观锁冲突:实例版本号不匹配",
137            ));
138        }
139
140        self.event_bus
141            .publish(WorkflowEvent::TransitionFired {
142                instance_id: instance_id.into(),
143                from: current_state.clone(),
144                to: transition.to.clone(),
145                event: event.into(),
146                timestamp: Utc::now(),
147            })
148            .await
149            .ok();
150
151        self.audit.log_transition(
152            "system",
153            instance_id,
154            "",
155            event,
156            InstanceStatus::Running,
157            InstanceStatus::Running,
158        );
159
160        Ok(TransitionResult {
161            from_state: current_state,
162            to_state: transition.to.clone(),
163            event: event.into(),
164            migrated: true,
165        })
166    }
167}
168
169#[cfg(test)]
170mod tests {
171    use super::*;
172    use crate::guard::DefaultGuardEvaluator;
173    use crate::instance::FlowInstance;
174    use crate::integration::SensitiveFieldRegistry;
175    use crate::observability::{AuditLogger, NoopEventBus};
176    use crate::repository::InMemoryInstanceRepository;
177
178    fn machine() -> StateMachineDefinition {
179        StateMachineDefinition {
180            initial_state: "draft".into(),
181            states: vec!["draft".into(), "review".into(), "done".into()],
182            transitions: vec![
183                crate::definition::Transition {
184                    from: "draft".into(),
185                    to: "review".into(),
186                    event: "submit".into(),
187                    guard: None,
188                },
189                crate::definition::Transition {
190                    from: "review".into(),
191                    to: "done".into(),
192                    event: "approve".into(),
193                    guard: Some("$.amount > 100".into()),
194                },
195            ],
196        }
197    }
198
199    fn setup() -> (StateMachineEngine, Arc<InMemoryInstanceRepository>) {
200        let repo = Arc::new(InMemoryInstanceRepository::default());
201        let engine = StateMachineEngine::new(
202            repo.clone(),
203            Arc::new(DefaultGuardEvaluator::default()),
204            Arc::new(NoopEventBus),
205            Arc::new(AuditLogger::new(Arc::new(SensitiveFieldRegistry::new()))),
206        );
207        (engine, repo)
208    }
209
210    #[tokio::test]
211    async fn fire_success() {
212        let (engine, repo) = setup();
213        let inst = FlowInstance::new(
214            "i1",
215            "test",
216            semver::Version::new(1, 0, 0),
217            "u1",
218            serde_json::json!({"current_state": "draft"}),
219            "start",
220        );
221        repo.create(&inst).await.unwrap();
222
223        let result = engine
224            .fire("i1", "submit", &machine(), serde_json::json!({}))
225            .await
226            .unwrap();
227        assert!(result.migrated);
228        assert_eq!(result.from_state, "draft");
229        assert_eq!(result.to_state, "review");
230    }
231
232    #[tokio::test]
233    async fn fire_no_matching_transition() {
234        let (engine, repo) = setup();
235        let inst = FlowInstance::new(
236            "i1",
237            "test",
238            semver::Version::new(1, 0, 0),
239            "u1",
240            serde_json::json!({"current_state": "draft"}),
241            "start",
242        );
243        repo.create(&inst).await.unwrap();
244
245        let result = engine
246            .fire("i1", "approve", &machine(), serde_json::json!({}))
247            .await;
248        assert!(result.is_err());
249        assert_eq!(
250            result.unwrap_err().code,
251            WorkflowErrorCode::NoMatchingTransition
252        );
253    }
254
255    #[tokio::test]
256    async fn fire_instance_not_found() {
257        let (engine, _) = setup();
258        let result = engine
259            .fire("nonexistent", "submit", &machine(), serde_json::json!({}))
260            .await;
261        assert!(result.is_err());
262        assert_eq!(
263            result.unwrap_err().code,
264            WorkflowErrorCode::InstanceNotFound
265        );
266    }
267
268    #[tokio::test]
269    async fn fire_guard_false_drops_event() {
270        let (engine, repo) = setup();
271        let inst = FlowInstance::new(
272            "i1",
273            "test",
274            semver::Version::new(1, 0, 0),
275            "u1",
276            serde_json::json!({"current_state": "review"}),
277            "start",
278        );
279        repo.create(&inst).await.unwrap();
280
281        let result = engine
282            .fire(
283                "i1",
284                "approve",
285                &machine(),
286                serde_json::json!({"amount": 50}),
287            )
288            .await
289            .unwrap();
290        assert!(!result.migrated);
291    }
292
293    #[tokio::test]
294    async fn fire_guard_true_migrates() {
295        let (engine, repo) = setup();
296        let inst = FlowInstance::new(
297            "i1",
298            "test",
299            semver::Version::new(1, 0, 0),
300            "u1",
301            serde_json::json!({"current_state": "review"}),
302            "start",
303        );
304        repo.create(&inst).await.unwrap();
305
306        let result = engine
307            .fire(
308                "i1",
309                "approve",
310                &machine(),
311                serde_json::json!({"amount": 200}),
312            )
313            .await
314            .unwrap();
315        assert!(result.migrated);
316        assert_eq!(result.to_state, "done");
317    }
318}