Skip to main content

tokio_prompt_orchestrator/
conversation_state.rs

1//! # Conversation State Machine
2//!
3//! Tracks the lifecycle phase of a conversation via a finite-state machine.
4//! Phases transition based on [`TransitionTrigger`] events detected from user
5//! text or explicit external signals (timeout, close).
6
7use std::collections::HashMap;
8
9// ── ConversationPhase ─────────────────────────────────────────────────────────
10
11/// High-level lifecycle phase of a conversation.
12#[derive(Debug, Clone, PartialEq, Eq, Hash)]
13pub enum ConversationPhase {
14    /// Conversation has just started; waiting for first meaningful input.
15    Opening,
16    /// User mentioned a task; gathering requirements.
17    TaskDiscovery,
18    /// Actively working on the user's task.
19    ActiveTask,
20    /// A clarifying question was asked; waiting for user reply.
21    Clarification,
22    /// Generating a summary / wrap-up response.
23    Summarizing,
24    /// Conversation ended gracefully.
25    Closing,
26    /// Conversation ended due to inactivity or abrupt user departure.
27    Abandoned,
28}
29
30impl ConversationPhase {
31    /// Returns `true` for phases from which no further transitions occur.
32    pub fn is_terminal(&self) -> bool {
33        matches!(self, Self::Closing | Self::Abandoned)
34    }
35}
36
37// ── TransitionTrigger ─────────────────────────────────────────────────────────
38
39/// Events that may cause a phase transition.
40#[derive(Debug, Clone, PartialEq, Eq, Hash)]
41pub enum TransitionTrigger {
42    /// User sent a greeting ("hello", "hi", …).
43    UserGreeting,
44    /// User mentioned a task or goal.
45    TaskMentioned,
46    /// User asked a question (heuristic: contains "?").
47    QuestionAsked,
48    /// The active task was completed successfully.
49    TaskCompleted,
50    /// User expressed satisfaction ("thank you", "thanks", …).
51    UserSatisfied,
52    /// An inactivity / session timeout fired.
53    TimeoutExpired,
54    /// User left or connection dropped.
55    UserLeft,
56    /// Explicit close signal from the system or user ("bye", "close", …).
57    ExplicitClose,
58}
59
60// ── PhaseTransition ───────────────────────────────────────────────────────────
61
62/// A recorded phase change with its cause and wall-clock timestamp.
63#[derive(Debug, Clone)]
64pub struct PhaseTransition {
65    /// Phase before the transition.
66    pub from: ConversationPhase,
67    /// Phase after the transition.
68    pub to: ConversationPhase,
69    /// What triggered this transition.
70    pub trigger: TransitionTrigger,
71    /// Unix-epoch milliseconds when the transition occurred.
72    pub timestamp: u64,
73}
74
75// ── ConversationStateMachine ──────────────────────────────────────────────────
76
77/// Finite-state machine that advances a conversation through its lifecycle.
78///
79/// # Example
80///
81/// ```rust
82/// use tokio_prompt_orchestrator::conversation_state::{
83///     ConversationStateMachine, TransitionTrigger,
84/// };
85///
86/// let mut sm = ConversationStateMachine::new();
87/// let new_phase = sm.trigger(TransitionTrigger::TaskMentioned, 1_000);
88/// assert!(new_phase.is_some());
89/// ```
90pub struct ConversationStateMachine {
91    phase: ConversationPhase,
92    phase_entered_at: u64,
93    history: Vec<PhaseTransition>,
94    /// (from, trigger) -> to  —  the rule table.
95    rules: HashMap<(ConversationPhase, TransitionTrigger), ConversationPhase>,
96}
97
98impl ConversationStateMachine {
99    /// Create a new state machine starting in [`ConversationPhase::Opening`].
100    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    /// Build the static rule table: (from, trigger) → to.
111    fn build_rules() -> HashMap<(ConversationPhase, TransitionTrigger), ConversationPhase> {
112        let mut r: HashMap<(ConversationPhase, TransitionTrigger), ConversationPhase> =
113            HashMap::new();
114
115        // Opening transitions
116        r.insert(
117            (ConversationPhase::Opening, TransitionTrigger::UserGreeting),
118            ConversationPhase::Opening, // greeting stays in opening
119        );
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        // TaskDiscovery transitions
138        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        // Clarification transitions
160        r.insert(
161            (ConversationPhase::Clarification, TransitionTrigger::TaskMentioned),
162            ConversationPhase::ActiveTask,
163        );
164        r.insert(
165            (ConversationPhase::Clarification, TransitionTrigger::QuestionAsked),
166            ConversationPhase::Clarification, // nested clarification
167        );
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        // ActiveTask transitions
182        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        // Summarizing transitions
204        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    /// Return the current phase.
229    pub fn current_phase(&self) -> &ConversationPhase {
230        &self.phase
231    }
232
233    /// Apply a trigger.  Returns the new phase if a transition occurred, or
234    /// `None` if the machine is in a terminal state or no rule matched.
235    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        // Record the transition (even self-loops, for history completeness).
247        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    /// Return `true` if the given trigger can fire from `from`.
260    pub fn can_transition(
261        from: &ConversationPhase,
262        trigger: &TransitionTrigger,
263    ) -> bool {
264        // Build a fresh instance to check without mutating self.
265        let rules = Self::build_rules();
266        rules.contains_key(&(from.clone(), trigger.clone()))
267    }
268
269    /// Return the full transition history.
270    pub fn history(&self) -> &[PhaseTransition] {
271        &self.history
272    }
273
274    /// Milliseconds spent in the current phase as of `now`.
275    pub fn time_in_current_phase(&self, now: u64) -> u64 {
276        now.saturating_sub(self.phase_entered_at)
277    }
278
279    /// Heuristic trigger detection from raw user text.
280    ///
281    /// Rules (checked in order):
282    /// - "bye" / "goodbye" / "close" → [`TransitionTrigger::ExplicitClose`]
283    /// - "thank" / "thanks" → [`TransitionTrigger::UserSatisfied`]
284    /// - "hello" / "hi" / "hey" → [`TransitionTrigger::UserGreeting`]
285    /// - text contains "?" → [`TransitionTrigger::QuestionAsked`]
286    /// - any other non-empty text → [`TransitionTrigger::TaskMentioned`]
287    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// ── Tests ─────────────────────────────────────────────────────────────────────
317
318#[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); // -> TaskDiscovery
364        sm.trigger(TransitionTrigger::TaskMentioned, 200); // -> ActiveTask
365        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}