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        }
272    }
273
274    // ─── PendingQueries tests ───
275
276    #[test]
277    fn pending_queries_single_feed() {
278        let mut pq = PendingQueries::new(vec![make_query(0)]);
279        assert_eq!(pq.remaining(), 1);
280        assert!(!pq.is_complete());
281
282        let complete = pq.feed(&QueryId::batch(0), "resp".into()).unwrap();
283        assert!(complete);
284        assert_eq!(pq.remaining(), 0);
285    }
286
287    #[test]
288    fn pending_queries_multi_feed_ordering() {
289        let mut pq = PendingQueries::new(vec![make_query(0), make_query(1), make_query(2)]);
290
291        // feed in reverse order
292        assert!(!pq.feed(&QueryId::batch(2), "resp-2".into()).unwrap());
293        assert!(!pq.feed(&QueryId::batch(0), "resp-0".into()).unwrap());
294        assert!(pq.feed(&QueryId::batch(1), "resp-1".into()).unwrap());
295
296        // into_ordered_responses returns in insertion order
297        let responses = pq.into_ordered_responses();
298        assert_eq!(responses, vec!["resp-0", "resp-1", "resp-2"]);
299    }
300
301    #[test]
302    fn pending_queries_unknown_query_error() {
303        let mut pq = PendingQueries::new(vec![make_query(0)]);
304        let err = pq.feed(&QueryId::batch(99), "resp".into()).unwrap_err();
305        assert!(matches!(err, FeedError::UnknownQuery(_)));
306    }
307
308    #[test]
309    fn pending_queries_double_feed_error() {
310        let mut pq = PendingQueries::new(vec![make_query(0)]);
311        pq.feed(&QueryId::batch(0), "resp".into()).unwrap();
312        let err = pq.feed(&QueryId::batch(0), "resp2".into()).unwrap_err();
313        assert!(matches!(err, FeedError::AlreadyResponded(_)));
314    }
315
316    #[test]
317    fn pending_queries_pending_list() {
318        let mut pq = PendingQueries::new(vec![make_query(0), make_query(1)]);
319        assert_eq!(pq.pending_queries().len(), 2);
320
321        pq.feed(&QueryId::batch(0), "resp".into()).unwrap();
322        let pending = pq.pending_queries();
323        assert_eq!(pending.len(), 1);
324        assert_eq!(pending[0].id, QueryId::batch(1));
325    }
326
327    #[test]
328    fn pending_queries_roundtrip_json() {
329        let mut pq = PendingQueries::new(vec![make_query(0), make_query(1)]);
330        pq.feed(&QueryId::batch(0), "resp-0".into()).unwrap();
331
332        let json = serde_json::to_value(&pq).unwrap();
333        let restored: PendingQueries = serde_json::from_value(json).unwrap();
334        assert_eq!(restored.remaining(), 1);
335        assert_eq!(restored.queries.len(), 2);
336    }
337
338    // ─── ExecutionState transition tests ───
339
340    #[test]
341    fn running_to_paused() {
342        let mut state = ExecutionState::Running;
343        state.pause(vec![make_query(0)]).unwrap();
344        assert_eq!(state.name(), "Paused");
345    }
346
347    #[test]
348    fn paused_feed_and_take() {
349        let mut state = ExecutionState::Running;
350        state.pause(vec![make_query(0), make_query(1)]).unwrap();
351
352        assert!(!state.feed(&QueryId::batch(0), "r0".into()).unwrap());
353        assert!(state.feed(&QueryId::batch(1), "r1".into()).unwrap());
354
355        let responses = state.take_responses().unwrap();
356        assert_eq!(responses, vec!["r0", "r1"]);
357        assert_eq!(state.name(), "Running");
358    }
359
360    #[test]
361    fn take_responses_incomplete_fails() {
362        let mut state = ExecutionState::Running;
363        state.pause(vec![make_query(0), make_query(1)]).unwrap();
364        state.feed(&QueryId::batch(0), "r0".into()).unwrap();
365
366        let err = state.take_responses().unwrap_err();
367        assert_eq!(err.actual, "Paused");
368        // state should remain Paused
369        assert_eq!(state.name(), "Paused");
370    }
371
372    #[test]
373    fn running_to_completed() {
374        let mut state = ExecutionState::Running;
375        state.complete(json!({"answer": 42})).unwrap();
376        assert!(state.is_terminal());
377        assert_eq!(state.name(), "Completed");
378    }
379
380    #[test]
381    fn running_to_failed() {
382        let mut state = ExecutionState::Running;
383        state.fail("boom".into()).unwrap();
384        assert!(state.is_terminal());
385        assert_eq!(state.name(), "Failed");
386    }
387
388    #[test]
389    fn cancel_from_running() {
390        let mut state = ExecutionState::Running;
391        state.cancel().unwrap();
392        assert!(state.is_terminal());
393        assert_eq!(state.name(), "Cancelled");
394    }
395
396    #[test]
397    fn cancel_from_paused() {
398        let mut state = ExecutionState::Running;
399        state.pause(vec![make_query(0)]).unwrap();
400        state.cancel().unwrap();
401        assert_eq!(state.name(), "Cancelled");
402    }
403
404    // ─── remaining() tests ───
405
406    #[test]
407    fn remaining_running_is_zero() {
408        let state = ExecutionState::Running;
409        assert_eq!(state.remaining(), 0);
410    }
411
412    #[test]
413    fn remaining_tracks_feeds() {
414        let mut state = ExecutionState::Running;
415        state
416            .pause(vec![make_query(0), make_query(1), make_query(2)])
417            .unwrap();
418        assert_eq!(state.remaining(), 3);
419
420        state.feed(&QueryId::batch(0), "r".into()).unwrap();
421        assert_eq!(state.remaining(), 2);
422
423        state.feed(&QueryId::batch(1), "r".into()).unwrap();
424        assert_eq!(state.remaining(), 1);
425    }
426
427    #[test]
428    fn remaining_terminal_is_zero() {
429        let state = ExecutionState::Completed {
430            result: json!(null),
431        };
432        assert_eq!(state.remaining(), 0);
433    }
434
435    // ─── Invalid transition tests ───
436
437    #[test]
438    fn feed_on_running_fails() {
439        let mut state = ExecutionState::Running;
440        let err = state.feed(&QueryId::single(), "r".into()).unwrap_err();
441        assert!(matches!(err, FeedError::InvalidState(_)));
442    }
443
444    #[test]
445    fn pause_on_paused_fails() {
446        let mut state = ExecutionState::Running;
447        state.pause(vec![make_query(0)]).unwrap();
448        let err = state.pause(vec![make_query(1)]).unwrap_err();
449        assert_eq!(err.expected, "Running");
450    }
451
452    #[test]
453    fn complete_on_paused_fails() {
454        let mut state = ExecutionState::Running;
455        state.pause(vec![make_query(0)]).unwrap();
456        let err = state.complete(json!(null)).unwrap_err();
457        assert_eq!(err.expected, "Running");
458    }
459
460    #[test]
461    fn cancel_on_completed_fails() {
462        let mut state = ExecutionState::Running;
463        state.complete(json!(null)).unwrap();
464        let err = state.cancel().unwrap_err();
465        assert_eq!(err.expected, "Running or Paused");
466    }
467
468    #[test]
469    fn cancel_on_failed_fails() {
470        let mut state = ExecutionState::Running;
471        state.fail("e".into()).unwrap();
472        let err = state.cancel().unwrap_err();
473        assert_eq!(err.expected, "Running or Paused");
474    }
475
476    #[test]
477    fn terminal_state_rejects_non_terminal() {
478        let state = ExecutionState::Running;
479        let err = TerminalState::try_from(state).unwrap_err();
480        assert_eq!(err.actual, "Running");
481    }
482
483    #[test]
484    fn terminal_state_from_completed() {
485        let state = ExecutionState::Completed { result: json!(42) };
486        let terminal = TerminalState::try_from(state).unwrap();
487        assert!(matches!(terminal, TerminalState::Completed { .. }));
488    }
489
490    #[test]
491    fn terminal_state_from_cancelled() {
492        let state = ExecutionState::Cancelled;
493        let terminal = TerminalState::try_from(state).unwrap();
494        assert!(matches!(terminal, TerminalState::Cancelled));
495    }
496}
497
498#[cfg(test)]
499mod proptests {
500    use super::*;
501    use crate::query::{LlmQuery, QueryId};
502    use proptest::prelude::*;
503
504    fn make_query(index: usize) -> LlmQuery {
505        LlmQuery {
506            id: QueryId::batch(index),
507            prompt: format!("prompt-{index}"),
508            system: None,
509            max_tokens: 100,
510            grounded: false,
511            underspecified: false,
512            cache_breakpoint: None,
513        }
514    }
515
516    proptest! {
517        /// into_ordered_responses returns insertion order regardless of feed order.
518        #[test]
519        fn feed_order_independent(size in 1usize..8) {
520            let queries: Vec<LlmQuery> = (0..size).map(make_query).collect();
521            let mut pq = PendingQueries::new(queries);
522
523            // feed in reverse order
524            for i in (0..size).rev() {
525                let _ = pq.feed(&QueryId::batch(i), format!("resp-{i}"));
526            }
527
528            let responses = pq.into_ordered_responses();
529            // must return in insertion order (0, 1, 2, ...)
530            for (i, resp) in responses.iter().enumerate() {
531                prop_assert_eq!(resp, &format!("resp-{i}"));
532            }
533        }
534
535        /// Feeding the same query twice returns AlreadyResponded error.
536        #[test]
537        fn double_feed_always_errors(size in 1usize..8, target in 0usize..8) {
538            let target = target % size; // clamp to valid range
539            let queries: Vec<LlmQuery> = (0..size).map(make_query).collect();
540            let mut pq = PendingQueries::new(queries);
541
542            pq.feed(&QueryId::batch(target), "first".into()).unwrap();
543            let err = pq.feed(&QueryId::batch(target), "second".into()).unwrap_err();
544            prop_assert!(matches!(err, FeedError::AlreadyResponded(_)));
545        }
546
547        /// Feeding a non-existent query_id returns UnknownQuery error.
548        #[test]
549        fn unknown_query_always_errors(size in 1usize..8, bad_id in 100usize..200) {
550            let queries: Vec<LlmQuery> = (0..size).map(make_query).collect();
551            let mut pq = PendingQueries::new(queries);
552
553            let err = pq.feed(&QueryId::batch(bad_id), "resp".into()).unwrap_err();
554            prop_assert!(matches!(err, FeedError::UnknownQuery(_)));
555        }
556
557        /// remaining() decreases by 1 with each feed.
558        #[test]
559        fn remaining_decreases_monotonically(size in 1usize..10) {
560            let queries: Vec<LlmQuery> = (0..size).map(make_query).collect();
561            let mut pq = PendingQueries::new(queries);
562
563            for i in 0..size {
564                prop_assert_eq!(pq.remaining(), size - i);
565                let _ = pq.feed(&QueryId::batch(i), format!("r-{i}"));
566            }
567            prop_assert_eq!(pq.remaining(), 0);
568            prop_assert!(pq.is_complete());
569        }
570    }
571}