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 }
272 }
273
274 #[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 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 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 #[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 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 #[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 #[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 #[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 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 for (i, resp) in responses.iter().enumerate() {
531 prop_assert_eq!(resp, &format!("resp-{i}"));
532 }
533 }
534
535 #[test]
537 fn double_feed_always_errors(size in 1usize..8, target in 0usize..8) {
538 let target = target % size; 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 #[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 #[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}