Skip to main content

algocline_core/
state.rs

1use std::collections::HashMap;
2
3use indexmap::IndexMap;
4use serde::{Deserialize, Serialize};
5
6use crate::query::{LlmQuery, QueryId};
7
8#[derive(Debug, thiserror::Error)]
9#[error("invalid state transition: expected {expected}, got {actual}")]
10pub struct TransitionError {
11    pub expected: &'static str,
12    pub actual: &'static str,
13}
14
15#[derive(Debug, thiserror::Error)]
16pub enum FeedError {
17    #[error("unknown query_id: {0}")]
18    UnknownQuery(QueryId),
19    #[error("already responded to query_id: {0}")]
20    AlreadyResponded(QueryId),
21    #[error(transparent)]
22    InvalidState(#[from] TransitionError),
23}
24
25/// Join barrier that collects N LLM responses.
26///
27/// Responses can be fed in any order and concurrency.
28/// Becomes complete when all queries have been responded to.
29#[derive(Debug, Serialize, Deserialize)]
30pub struct PendingQueries {
31    /// Issued queries (insertion order preserved via IndexMap).
32    queries: IndexMap<QueryId, LlmQuery>,
33    responses: HashMap<QueryId, String>,
34}
35
36impl PendingQueries {
37    pub fn new(queries: Vec<LlmQuery>) -> Self {
38        let map = queries
39            .into_iter()
40            .map(|q| (q.id.clone(), q))
41            .collect::<IndexMap<_, _>>();
42        Self {
43            queries: map,
44            responses: HashMap::new(),
45        }
46    }
47
48    /// Feed one response. Returns `true` if all queries are now complete.
49    pub fn feed(&mut self, id: &QueryId, response: String) -> Result<bool, FeedError> {
50        if !self.queries.contains_key(id) {
51            return Err(FeedError::UnknownQuery(id.clone()));
52        }
53        if self.responses.contains_key(id) {
54            return Err(FeedError::AlreadyResponded(id.clone()));
55        }
56        self.responses.insert(id.clone(), response);
57        Ok(self.is_complete())
58    }
59
60    pub fn pending_queries(&self) -> Vec<&LlmQuery> {
61        self.queries
62            .values()
63            .filter(|q| !self.responses.contains_key(&q.id))
64            .collect()
65    }
66
67    pub fn remaining(&self) -> usize {
68        self.queries.len() - self.responses.len()
69    }
70
71    pub fn is_complete(&self) -> bool {
72        self.responses.len() == self.queries.len()
73    }
74
75    /// Consume and return responses in query insertion order.
76    /// Corresponds to the Paused → Running transition.
77    pub fn into_ordered_responses(self) -> Vec<String> {
78        self.queries
79            .keys()
80            .map(|id| {
81                // is_complete() guarantees queries and responses share the same key set,
82                // but fall back to empty string if called without checking is_complete()
83                self.responses.get(id).cloned().unwrap_or_default()
84            })
85            .collect()
86    }
87}
88
89pub enum ExecutionState {
90    Running,
91    /// Awaiting 1..N LLM responses.
92    Paused(PendingQueries),
93    Completed {
94        result: serde_json::Value,
95    },
96    Failed {
97        error: String,
98    },
99    /// Explicit cancellation by the host.
100    Cancelled,
101}
102
103impl ExecutionState {
104    pub fn is_terminal(&self) -> bool {
105        matches!(
106            self,
107            Self::Completed { .. } | Self::Failed { .. } | Self::Cancelled
108        )
109    }
110
111    /// Number of pending queries. Returns 0 for non-Paused states.
112    pub fn remaining(&self) -> usize {
113        match self {
114            Self::Paused(pending) => pending.remaining(),
115            _ => 0,
116        }
117    }
118
119    /// Returns the state name (for error messages).
120    pub fn name(&self) -> &'static str {
121        match self {
122            Self::Running => "Running",
123            Self::Paused(_) => "Paused",
124            Self::Completed { .. } => "Completed",
125            Self::Failed { .. } => "Failed",
126            Self::Cancelled => "Cancelled",
127        }
128    }
129
130    /// Feed a response. Only valid in Paused state.
131    /// Returns `Ok(true)` when all queries are complete, `Ok(false)` otherwise.
132    pub fn feed(&mut self, id: &QueryId, response: String) -> Result<bool, FeedError> {
133        match self {
134            Self::Paused(pending) => pending.feed(id, response),
135            other => Err(TransitionError {
136                expected: "Paused",
137                actual: other.name(),
138            }
139            .into()),
140        }
141    }
142
143    /// Extract responses from a complete Paused state.
144    /// Transitions self to Running (preparing for Lua resumption).
145    pub fn take_responses(&mut self) -> Result<Vec<String>, TransitionError> {
146        match std::mem::replace(self, Self::Running) {
147            Self::Paused(pending) if pending.is_complete() => Ok(pending.into_ordered_responses()),
148            prev => {
149                let actual = prev.name();
150                *self = prev;
151                Err(TransitionError {
152                    expected: "Paused(complete)",
153                    actual,
154                })
155            }
156        }
157    }
158
159    /// Running → Completed.
160    pub fn complete(&mut self, result: serde_json::Value) -> Result<(), TransitionError> {
161        match self {
162            Self::Running => {
163                *self = Self::Completed { result };
164                Ok(())
165            }
166            other => Err(TransitionError {
167                expected: "Running",
168                actual: other.name(),
169            }),
170        }
171    }
172
173    /// Running → Failed.
174    pub fn fail(&mut self, error: String) -> Result<(), TransitionError> {
175        match self {
176            Self::Running => {
177                *self = Self::Failed { error };
178                Ok(())
179            }
180            other => Err(TransitionError {
181                expected: "Running",
182                actual: other.name(),
183            }),
184        }
185    }
186
187    /// Running → Paused (triggered by alc.llm() / alc.llm_batch()).
188    pub fn pause(&mut self, queries: Vec<LlmQuery>) -> Result<(), TransitionError> {
189        match self {
190            Self::Running => {
191                *self = Self::Paused(PendingQueries::new(queries));
192                Ok(())
193            }
194            other => Err(TransitionError {
195                expected: "Running",
196                actual: other.name(),
197            }),
198        }
199    }
200
201    /// Running | Paused → Cancelled (explicit host cancellation).
202    pub fn cancel(&mut self) -> Result<(), TransitionError> {
203        match self {
204            Self::Running | Self::Paused(_) => {
205                *self = Self::Cancelled;
206                Ok(())
207            }
208            other => Err(TransitionError {
209                expected: "Running or Paused",
210                actual: other.name(),
211            }),
212        }
213    }
214}
215
216/// Return type of Session.resume(). Never returns Running.
217pub enum ResumeOutcome {
218    /// Lua resumed and paused again at alc.llm().
219    Paused {
220        queries: Vec<LlmQuery>,
221    },
222    Completed {
223        result: serde_json::Value,
224    },
225    Failed {
226        error: String,
227    },
228    /// Cancelled during resume.
229    Cancelled,
230}
231
232/// Terminal execution state. Only Completed, Failed, or Cancelled.
233#[derive(Debug, serde::Serialize, serde::Deserialize)]
234pub enum TerminalState {
235    Completed { result: serde_json::Value },
236    Failed { error: String },
237    Cancelled,
238}
239
240impl TryFrom<ExecutionState> for TerminalState {
241    type Error = TransitionError;
242
243    fn try_from(state: ExecutionState) -> Result<Self, TransitionError> {
244        match state {
245            ExecutionState::Completed { result } => Ok(Self::Completed { result }),
246            ExecutionState::Failed { error } => Ok(Self::Failed { error }),
247            ExecutionState::Cancelled => Ok(Self::Cancelled),
248            other => Err(TransitionError {
249                expected: "Completed, Failed, or Cancelled",
250                actual: other.name(),
251            }),
252        }
253    }
254}
255
256#[cfg(test)]
257mod tests {
258    use super::*;
259    use crate::query::{LlmQuery, QueryId};
260    use serde_json::json;
261
262    fn make_query(index: usize) -> LlmQuery {
263        LlmQuery {
264            id: QueryId::batch(index),
265            prompt: format!("prompt-{index}"),
266            system: None,
267            max_tokens: 100,
268            grounded: false,
269            underspecified: false,
270            cache_breakpoint: None,
271            role: None,
272        }
273    }
274
275    // ─── PendingQueries tests ───
276
277    #[test]
278    fn pending_queries_single_feed() {
279        let mut pq = PendingQueries::new(vec![make_query(0)]);
280        assert_eq!(pq.remaining(), 1);
281        assert!(!pq.is_complete());
282
283        let complete = pq.feed(&QueryId::batch(0), "resp".into()).unwrap();
284        assert!(complete);
285        assert_eq!(pq.remaining(), 0);
286    }
287
288    #[test]
289    fn pending_queries_multi_feed_ordering() {
290        let mut pq = PendingQueries::new(vec![make_query(0), make_query(1), make_query(2)]);
291
292        // feed in reverse order
293        assert!(!pq.feed(&QueryId::batch(2), "resp-2".into()).unwrap());
294        assert!(!pq.feed(&QueryId::batch(0), "resp-0".into()).unwrap());
295        assert!(pq.feed(&QueryId::batch(1), "resp-1".into()).unwrap());
296
297        // into_ordered_responses returns in insertion order
298        let responses = pq.into_ordered_responses();
299        assert_eq!(responses, vec!["resp-0", "resp-1", "resp-2"]);
300    }
301
302    #[test]
303    fn pending_queries_unknown_query_error() {
304        let mut pq = PendingQueries::new(vec![make_query(0)]);
305        let err = pq.feed(&QueryId::batch(99), "resp".into()).unwrap_err();
306        assert!(matches!(err, FeedError::UnknownQuery(_)));
307    }
308
309    #[test]
310    fn pending_queries_double_feed_error() {
311        let mut pq = PendingQueries::new(vec![make_query(0)]);
312        pq.feed(&QueryId::batch(0), "resp".into()).unwrap();
313        let err = pq.feed(&QueryId::batch(0), "resp2".into()).unwrap_err();
314        assert!(matches!(err, FeedError::AlreadyResponded(_)));
315    }
316
317    #[test]
318    fn pending_queries_pending_list() {
319        let mut pq = PendingQueries::new(vec![make_query(0), make_query(1)]);
320        assert_eq!(pq.pending_queries().len(), 2);
321
322        pq.feed(&QueryId::batch(0), "resp".into()).unwrap();
323        let pending = pq.pending_queries();
324        assert_eq!(pending.len(), 1);
325        assert_eq!(pending[0].id, QueryId::batch(1));
326    }
327
328    #[test]
329    fn pending_queries_roundtrip_json() {
330        let mut pq = PendingQueries::new(vec![make_query(0), make_query(1)]);
331        pq.feed(&QueryId::batch(0), "resp-0".into()).unwrap();
332
333        let json = serde_json::to_value(&pq).unwrap();
334        let restored: PendingQueries = serde_json::from_value(json).unwrap();
335        assert_eq!(restored.remaining(), 1);
336        assert_eq!(restored.queries.len(), 2);
337    }
338
339    // ─── ExecutionState transition tests ───
340
341    #[test]
342    fn running_to_paused() {
343        let mut state = ExecutionState::Running;
344        state.pause(vec![make_query(0)]).unwrap();
345        assert_eq!(state.name(), "Paused");
346    }
347
348    #[test]
349    fn paused_feed_and_take() {
350        let mut state = ExecutionState::Running;
351        state.pause(vec![make_query(0), make_query(1)]).unwrap();
352
353        assert!(!state.feed(&QueryId::batch(0), "r0".into()).unwrap());
354        assert!(state.feed(&QueryId::batch(1), "r1".into()).unwrap());
355
356        let responses = state.take_responses().unwrap();
357        assert_eq!(responses, vec!["r0", "r1"]);
358        assert_eq!(state.name(), "Running");
359    }
360
361    #[test]
362    fn take_responses_incomplete_fails() {
363        let mut state = ExecutionState::Running;
364        state.pause(vec![make_query(0), make_query(1)]).unwrap();
365        state.feed(&QueryId::batch(0), "r0".into()).unwrap();
366
367        let err = state.take_responses().unwrap_err();
368        assert_eq!(err.actual, "Paused");
369        // state should remain Paused
370        assert_eq!(state.name(), "Paused");
371    }
372
373    #[test]
374    fn running_to_completed() {
375        let mut state = ExecutionState::Running;
376        state.complete(json!({"answer": 42})).unwrap();
377        assert!(state.is_terminal());
378        assert_eq!(state.name(), "Completed");
379    }
380
381    #[test]
382    fn running_to_failed() {
383        let mut state = ExecutionState::Running;
384        state.fail("boom".into()).unwrap();
385        assert!(state.is_terminal());
386        assert_eq!(state.name(), "Failed");
387    }
388
389    #[test]
390    fn cancel_from_running() {
391        let mut state = ExecutionState::Running;
392        state.cancel().unwrap();
393        assert!(state.is_terminal());
394        assert_eq!(state.name(), "Cancelled");
395    }
396
397    #[test]
398    fn cancel_from_paused() {
399        let mut state = ExecutionState::Running;
400        state.pause(vec![make_query(0)]).unwrap();
401        state.cancel().unwrap();
402        assert_eq!(state.name(), "Cancelled");
403    }
404
405    // ─── remaining() tests ───
406
407    #[test]
408    fn remaining_running_is_zero() {
409        let state = ExecutionState::Running;
410        assert_eq!(state.remaining(), 0);
411    }
412
413    #[test]
414    fn remaining_tracks_feeds() {
415        let mut state = ExecutionState::Running;
416        state
417            .pause(vec![make_query(0), make_query(1), make_query(2)])
418            .unwrap();
419        assert_eq!(state.remaining(), 3);
420
421        state.feed(&QueryId::batch(0), "r".into()).unwrap();
422        assert_eq!(state.remaining(), 2);
423
424        state.feed(&QueryId::batch(1), "r".into()).unwrap();
425        assert_eq!(state.remaining(), 1);
426    }
427
428    #[test]
429    fn remaining_terminal_is_zero() {
430        let state = ExecutionState::Completed {
431            result: json!(null),
432        };
433        assert_eq!(state.remaining(), 0);
434    }
435
436    // ─── Invalid transition tests ───
437
438    #[test]
439    fn feed_on_running_fails() {
440        let mut state = ExecutionState::Running;
441        let err = state.feed(&QueryId::single(), "r".into()).unwrap_err();
442        assert!(matches!(err, FeedError::InvalidState(_)));
443    }
444
445    #[test]
446    fn pause_on_paused_fails() {
447        let mut state = ExecutionState::Running;
448        state.pause(vec![make_query(0)]).unwrap();
449        let err = state.pause(vec![make_query(1)]).unwrap_err();
450        assert_eq!(err.expected, "Running");
451    }
452
453    #[test]
454    fn complete_on_paused_fails() {
455        let mut state = ExecutionState::Running;
456        state.pause(vec![make_query(0)]).unwrap();
457        let err = state.complete(json!(null)).unwrap_err();
458        assert_eq!(err.expected, "Running");
459    }
460
461    #[test]
462    fn cancel_on_completed_fails() {
463        let mut state = ExecutionState::Running;
464        state.complete(json!(null)).unwrap();
465        let err = state.cancel().unwrap_err();
466        assert_eq!(err.expected, "Running or Paused");
467    }
468
469    #[test]
470    fn cancel_on_failed_fails() {
471        let mut state = ExecutionState::Running;
472        state.fail("e".into()).unwrap();
473        let err = state.cancel().unwrap_err();
474        assert_eq!(err.expected, "Running or Paused");
475    }
476
477    #[test]
478    fn terminal_state_rejects_non_terminal() {
479        let state = ExecutionState::Running;
480        let err = TerminalState::try_from(state).unwrap_err();
481        assert_eq!(err.actual, "Running");
482    }
483
484    #[test]
485    fn terminal_state_from_completed() {
486        let state = ExecutionState::Completed { result: json!(42) };
487        let terminal = TerminalState::try_from(state).unwrap();
488        assert!(matches!(terminal, TerminalState::Completed { .. }));
489    }
490
491    #[test]
492    fn terminal_state_from_cancelled() {
493        let state = ExecutionState::Cancelled;
494        let terminal = TerminalState::try_from(state).unwrap();
495        assert!(matches!(terminal, TerminalState::Cancelled));
496    }
497}
498
499#[cfg(test)]
500mod proptests {
501    use super::*;
502    use crate::query::{LlmQuery, QueryId};
503    use proptest::prelude::*;
504
505    fn make_query(index: usize) -> LlmQuery {
506        LlmQuery {
507            id: QueryId::batch(index),
508            prompt: format!("prompt-{index}"),
509            system: None,
510            max_tokens: 100,
511            grounded: false,
512            underspecified: false,
513            cache_breakpoint: None,
514            role: None,
515        }
516    }
517
518    proptest! {
519        /// into_ordered_responses returns insertion order regardless of feed order.
520        #[test]
521        fn feed_order_independent(size in 1usize..8) {
522            let queries: Vec<LlmQuery> = (0..size).map(make_query).collect();
523            let mut pq = PendingQueries::new(queries);
524
525            // feed in reverse order
526            for i in (0..size).rev() {
527                let _ = pq.feed(&QueryId::batch(i), format!("resp-{i}"));
528            }
529
530            let responses = pq.into_ordered_responses();
531            // must return in insertion order (0, 1, 2, ...)
532            for (i, resp) in responses.iter().enumerate() {
533                prop_assert_eq!(resp, &format!("resp-{i}"));
534            }
535        }
536
537        /// Feeding the same query twice returns AlreadyResponded error.
538        #[test]
539        fn double_feed_always_errors(size in 1usize..8, target in 0usize..8) {
540            let target = target % size; // clamp to valid range
541            let queries: Vec<LlmQuery> = (0..size).map(make_query).collect();
542            let mut pq = PendingQueries::new(queries);
543
544            pq.feed(&QueryId::batch(target), "first".into()).unwrap();
545            let err = pq.feed(&QueryId::batch(target), "second".into()).unwrap_err();
546            prop_assert!(matches!(err, FeedError::AlreadyResponded(_)));
547        }
548
549        /// Feeding a non-existent query_id returns UnknownQuery error.
550        #[test]
551        fn unknown_query_always_errors(size in 1usize..8, bad_id in 100usize..200) {
552            let queries: Vec<LlmQuery> = (0..size).map(make_query).collect();
553            let mut pq = PendingQueries::new(queries);
554
555            let err = pq.feed(&QueryId::batch(bad_id), "resp".into()).unwrap_err();
556            prop_assert!(matches!(err, FeedError::UnknownQuery(_)));
557        }
558
559        /// remaining() decreases by 1 with each feed.
560        #[test]
561        fn remaining_decreases_monotonically(size in 1usize..10) {
562            let queries: Vec<LlmQuery> = (0..size).map(make_query).collect();
563            let mut pq = PendingQueries::new(queries);
564
565            for i in 0..size {
566                prop_assert_eq!(pq.remaining(), size - i);
567                let _ = pq.feed(&QueryId::batch(i), format!("r-{i}"));
568            }
569            prop_assert_eq!(pq.remaining(), 0);
570            prop_assert!(pq.is_complete());
571        }
572    }
573}