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#[derive(Debug, Serialize, Deserialize)]
30pub struct PendingQueries {
31 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 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 pub fn into_ordered_responses(self) -> Vec<String> {
78 self.queries
79 .keys()
80 .map(|id| {
81 self.responses.get(id).cloned().unwrap_or_default()
84 })
85 .collect()
86 }
87}
88
89pub enum ExecutionState {
90 Running,
91 Paused(PendingQueries),
93 Completed {
94 result: serde_json::Value,
95 },
96 Failed {
97 error: String,
98 },
99 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 pub fn remaining(&self) -> usize {
113 match self {
114 Self::Paused(pending) => pending.remaining(),
115 _ => 0,
116 }
117 }
118
119 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 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 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 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 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 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 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
216pub enum ResumeOutcome {
218 Paused {
220 queries: Vec<LlmQuery>,
221 },
222 Completed {
223 result: serde_json::Value,
224 },
225 Failed {
226 error: String,
227 },
228 Cancelled,
230}
231
232#[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 #[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 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 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 #[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 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 #[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 #[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 #[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 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 for (i, resp) in responses.iter().enumerate() {
533 prop_assert_eq!(resp, &format!("resp-{i}"));
534 }
535 }
536
537 #[test]
539 fn double_feed_always_errors(size in 1usize..8, target in 0usize..8) {
540 let target = target % size; 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 #[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 #[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}