1use std::collections::HashMap;
8
9#[derive(Debug, Clone, PartialEq, Eq, Hash)]
13pub enum ConversationPhase {
14 Opening,
16 TaskDiscovery,
18 ActiveTask,
20 Clarification,
22 Summarizing,
24 Closing,
26 Abandoned,
28}
29
30impl ConversationPhase {
31 pub fn is_terminal(&self) -> bool {
33 matches!(self, Self::Closing | Self::Abandoned)
34 }
35}
36
37#[derive(Debug, Clone, PartialEq, Eq, Hash)]
41pub enum TransitionTrigger {
42 UserGreeting,
44 TaskMentioned,
46 QuestionAsked,
48 TaskCompleted,
50 UserSatisfied,
52 TimeoutExpired,
54 UserLeft,
56 ExplicitClose,
58}
59
60#[derive(Debug, Clone)]
64pub struct PhaseTransition {
65 pub from: ConversationPhase,
67 pub to: ConversationPhase,
69 pub trigger: TransitionTrigger,
71 pub timestamp: u64,
73}
74
75pub struct ConversationStateMachine {
91 phase: ConversationPhase,
92 phase_entered_at: u64,
93 history: Vec<PhaseTransition>,
94 rules: HashMap<(ConversationPhase, TransitionTrigger), ConversationPhase>,
96}
97
98impl ConversationStateMachine {
99 pub fn new() -> Self {
101 let rules = Self::build_rules();
102 Self {
103 phase: ConversationPhase::Opening,
104 phase_entered_at: 0,
105 history: Vec::new(),
106 rules,
107 }
108 }
109
110 fn build_rules() -> HashMap<(ConversationPhase, TransitionTrigger), ConversationPhase> {
112 let mut r: HashMap<(ConversationPhase, TransitionTrigger), ConversationPhase> =
113 HashMap::new();
114
115 r.insert(
117 (ConversationPhase::Opening, TransitionTrigger::UserGreeting),
118 ConversationPhase::Opening, );
120 r.insert(
121 (ConversationPhase::Opening, TransitionTrigger::TaskMentioned),
122 ConversationPhase::TaskDiscovery,
123 );
124 r.insert(
125 (ConversationPhase::Opening, TransitionTrigger::TimeoutExpired),
126 ConversationPhase::Abandoned,
127 );
128 r.insert(
129 (ConversationPhase::Opening, TransitionTrigger::UserLeft),
130 ConversationPhase::Abandoned,
131 );
132 r.insert(
133 (ConversationPhase::Opening, TransitionTrigger::ExplicitClose),
134 ConversationPhase::Closing,
135 );
136
137 r.insert(
139 (ConversationPhase::TaskDiscovery, TransitionTrigger::QuestionAsked),
140 ConversationPhase::Clarification,
141 );
142 r.insert(
143 (ConversationPhase::TaskDiscovery, TransitionTrigger::TaskMentioned),
144 ConversationPhase::ActiveTask,
145 );
146 r.insert(
147 (ConversationPhase::TaskDiscovery, TransitionTrigger::TimeoutExpired),
148 ConversationPhase::Abandoned,
149 );
150 r.insert(
151 (ConversationPhase::TaskDiscovery, TransitionTrigger::UserLeft),
152 ConversationPhase::Abandoned,
153 );
154 r.insert(
155 (ConversationPhase::TaskDiscovery, TransitionTrigger::ExplicitClose),
156 ConversationPhase::Closing,
157 );
158
159 r.insert(
161 (ConversationPhase::Clarification, TransitionTrigger::TaskMentioned),
162 ConversationPhase::ActiveTask,
163 );
164 r.insert(
165 (ConversationPhase::Clarification, TransitionTrigger::QuestionAsked),
166 ConversationPhase::Clarification, );
168 r.insert(
169 (ConversationPhase::Clarification, TransitionTrigger::TimeoutExpired),
170 ConversationPhase::Abandoned,
171 );
172 r.insert(
173 (ConversationPhase::Clarification, TransitionTrigger::UserLeft),
174 ConversationPhase::Abandoned,
175 );
176 r.insert(
177 (ConversationPhase::Clarification, TransitionTrigger::ExplicitClose),
178 ConversationPhase::Closing,
179 );
180
181 r.insert(
183 (ConversationPhase::ActiveTask, TransitionTrigger::TaskCompleted),
184 ConversationPhase::Summarizing,
185 );
186 r.insert(
187 (ConversationPhase::ActiveTask, TransitionTrigger::QuestionAsked),
188 ConversationPhase::Clarification,
189 );
190 r.insert(
191 (ConversationPhase::ActiveTask, TransitionTrigger::TimeoutExpired),
192 ConversationPhase::Abandoned,
193 );
194 r.insert(
195 (ConversationPhase::ActiveTask, TransitionTrigger::UserLeft),
196 ConversationPhase::Abandoned,
197 );
198 r.insert(
199 (ConversationPhase::ActiveTask, TransitionTrigger::ExplicitClose),
200 ConversationPhase::Closing,
201 );
202
203 r.insert(
205 (ConversationPhase::Summarizing, TransitionTrigger::UserSatisfied),
206 ConversationPhase::Closing,
207 );
208 r.insert(
209 (ConversationPhase::Summarizing, TransitionTrigger::TaskMentioned),
210 ConversationPhase::TaskDiscovery,
211 );
212 r.insert(
213 (ConversationPhase::Summarizing, TransitionTrigger::ExplicitClose),
214 ConversationPhase::Closing,
215 );
216 r.insert(
217 (ConversationPhase::Summarizing, TransitionTrigger::TimeoutExpired),
218 ConversationPhase::Abandoned,
219 );
220 r.insert(
221 (ConversationPhase::Summarizing, TransitionTrigger::UserLeft),
222 ConversationPhase::Abandoned,
223 );
224
225 r
226 }
227
228 pub fn current_phase(&self) -> &ConversationPhase {
230 &self.phase
231 }
232
233 pub fn trigger(
236 &mut self,
237 trigger: TransitionTrigger,
238 now: u64,
239 ) -> Option<ConversationPhase> {
240 if self.phase.is_terminal() {
241 return None;
242 }
243 let key = (self.phase.clone(), trigger.clone());
244 let next = self.rules.get(&key)?.clone();
245
246 self.history.push(PhaseTransition {
248 from: self.phase.clone(),
249 to: next.clone(),
250 trigger,
251 timestamp: now,
252 });
253
254 self.phase = next.clone();
255 self.phase_entered_at = now;
256 Some(next)
257 }
258
259 pub fn can_transition(
261 from: &ConversationPhase,
262 trigger: &TransitionTrigger,
263 ) -> bool {
264 let rules = Self::build_rules();
266 rules.contains_key(&(from.clone(), trigger.clone()))
267 }
268
269 pub fn history(&self) -> &[PhaseTransition] {
271 &self.history
272 }
273
274 pub fn time_in_current_phase(&self, now: u64) -> u64 {
276 now.saturating_sub(self.phase_entered_at)
277 }
278
279 pub fn detect_trigger(user_text: &str) -> Option<TransitionTrigger> {
288 let lower = user_text.to_lowercase();
289 if lower.is_empty() {
290 return None;
291 }
292 if lower.contains("bye") || lower.contains("goodbye") || lower.contains("close") {
293 return Some(TransitionTrigger::ExplicitClose);
294 }
295 if lower.contains("thank") {
296 return Some(TransitionTrigger::UserSatisfied);
297 }
298 if lower.contains("hello") || lower.contains("hi ") || lower.starts_with("hi")
299 || lower.contains("hey")
300 {
301 return Some(TransitionTrigger::UserGreeting);
302 }
303 if lower.contains('?') {
304 return Some(TransitionTrigger::QuestionAsked);
305 }
306 Some(TransitionTrigger::TaskMentioned)
307 }
308}
309
310impl Default for ConversationStateMachine {
311 fn default() -> Self {
312 Self::new()
313 }
314}
315
316#[cfg(test)]
319mod tests {
320 use super::*;
321
322 #[test]
323 fn test_initial_phase_is_opening() {
324 let sm = ConversationStateMachine::new();
325 assert_eq!(sm.current_phase(), &ConversationPhase::Opening);
326 }
327
328 #[test]
329 fn test_terminal_phases() {
330 assert!(ConversationPhase::Closing.is_terminal());
331 assert!(ConversationPhase::Abandoned.is_terminal());
332 assert!(!ConversationPhase::Opening.is_terminal());
333 assert!(!ConversationPhase::ActiveTask.is_terminal());
334 }
335
336 #[test]
337 fn test_greeting_stays_in_opening() {
338 let mut sm = ConversationStateMachine::new();
339 let result = sm.trigger(TransitionTrigger::UserGreeting, 100);
340 assert!(result.is_some());
341 assert_eq!(sm.current_phase(), &ConversationPhase::Opening);
342 }
343
344 #[test]
345 fn test_task_mentioned_from_opening_goes_to_discovery() {
346 let mut sm = ConversationStateMachine::new();
347 let result = sm.trigger(TransitionTrigger::TaskMentioned, 200);
348 assert_eq!(result, Some(ConversationPhase::TaskDiscovery));
349 assert_eq!(sm.current_phase(), &ConversationPhase::TaskDiscovery);
350 }
351
352 #[test]
353 fn test_clarification_from_discovery() {
354 let mut sm = ConversationStateMachine::new();
355 sm.trigger(TransitionTrigger::TaskMentioned, 100);
356 let result = sm.trigger(TransitionTrigger::QuestionAsked, 200);
357 assert_eq!(result, Some(ConversationPhase::Clarification));
358 }
359
360 #[test]
361 fn test_active_task_to_summarizing() {
362 let mut sm = ConversationStateMachine::new();
363 sm.trigger(TransitionTrigger::TaskMentioned, 100); sm.trigger(TransitionTrigger::TaskMentioned, 200); let result = sm.trigger(TransitionTrigger::TaskCompleted, 300);
366 assert_eq!(result, Some(ConversationPhase::Summarizing));
367 }
368
369 #[test]
370 fn test_satisfied_in_summarizing_goes_to_closing() {
371 let mut sm = ConversationStateMachine::new();
372 sm.trigger(TransitionTrigger::TaskMentioned, 100);
373 sm.trigger(TransitionTrigger::TaskMentioned, 200);
374 sm.trigger(TransitionTrigger::TaskCompleted, 300);
375 let result = sm.trigger(TransitionTrigger::UserSatisfied, 400);
376 assert_eq!(result, Some(ConversationPhase::Closing));
377 assert!(sm.current_phase().is_terminal());
378 }
379
380 #[test]
381 fn test_no_transition_from_terminal() {
382 let mut sm = ConversationStateMachine::new();
383 sm.trigger(TransitionTrigger::ExplicitClose, 100);
384 assert_eq!(sm.current_phase(), &ConversationPhase::Closing);
385 let result = sm.trigger(TransitionTrigger::TaskMentioned, 200);
386 assert!(result.is_none());
387 }
388
389 #[test]
390 fn test_timeout_leads_to_abandoned() {
391 let mut sm = ConversationStateMachine::new();
392 let result = sm.trigger(TransitionTrigger::TimeoutExpired, 500);
393 assert_eq!(result, Some(ConversationPhase::Abandoned));
394 assert!(sm.current_phase().is_terminal());
395 }
396
397 #[test]
398 fn test_history_recorded() {
399 let mut sm = ConversationStateMachine::new();
400 sm.trigger(TransitionTrigger::TaskMentioned, 100);
401 sm.trigger(TransitionTrigger::QuestionAsked, 200);
402 assert_eq!(sm.history().len(), 2);
403 assert_eq!(sm.history()[0].trigger, TransitionTrigger::TaskMentioned);
404 assert_eq!(sm.history()[1].trigger, TransitionTrigger::QuestionAsked);
405 }
406
407 #[test]
408 fn test_time_in_current_phase() {
409 let mut sm = ConversationStateMachine::new();
410 sm.trigger(TransitionTrigger::TaskMentioned, 1000);
411 assert_eq!(sm.time_in_current_phase(1500), 500);
412 }
413
414 #[test]
415 fn test_can_transition() {
416 assert!(ConversationStateMachine::can_transition(
417 &ConversationPhase::Opening,
418 &TransitionTrigger::TaskMentioned
419 ));
420 assert!(!ConversationStateMachine::can_transition(
421 &ConversationPhase::Closing,
422 &TransitionTrigger::TaskMentioned
423 ));
424 }
425
426 #[test]
427 fn test_detect_trigger_greeting() {
428 assert_eq!(
429 ConversationStateMachine::detect_trigger("Hello there!"),
430 Some(TransitionTrigger::UserGreeting)
431 );
432 }
433
434 #[test]
435 fn test_detect_trigger_question() {
436 assert_eq!(
437 ConversationStateMachine::detect_trigger("Can you help me?"),
438 Some(TransitionTrigger::QuestionAsked)
439 );
440 }
441
442 #[test]
443 fn test_detect_trigger_satisfied() {
444 assert_eq!(
445 ConversationStateMachine::detect_trigger("Thank you so much!"),
446 Some(TransitionTrigger::UserSatisfied)
447 );
448 }
449
450 #[test]
451 fn test_detect_trigger_close() {
452 assert_eq!(
453 ConversationStateMachine::detect_trigger("Goodbye!"),
454 Some(TransitionTrigger::ExplicitClose)
455 );
456 }
457
458 #[test]
459 fn test_detect_trigger_task() {
460 assert_eq!(
461 ConversationStateMachine::detect_trigger("I need to process 1000 CSV files."),
462 Some(TransitionTrigger::TaskMentioned)
463 );
464 }
465
466 #[test]
467 fn test_detect_trigger_empty() {
468 assert_eq!(ConversationStateMachine::detect_trigger(""), None);
469 }
470}