Skip to main content

kcode_kennedy_stepped_turn_runtime/
lib.rs

1pub use kcode_kennedy_sessions::PendingTurnAdmission;
2use kcode_kennedy_sessions::{Session, TurnBoundary, TurnDeadline};
3use serde_json::Value;
4use std::collections::{HashMap, VecDeque};
5use std::future::Future;
6use std::sync::{Arc, Mutex, MutexGuard};
7use tokio::sync::Notify;
8use uuid::Uuid;
9
10pub struct QueuedAdmission<D> {
11    pub key: String,
12    pub recorded_at: String,
13    pub admission: PendingTurnAdmission,
14    pub delivery: D,
15}
16
17#[derive(Clone, Copy, Debug, Eq, PartialEq)]
18pub enum PushResult {
19    Queued { sequence: u64 },
20    Duplicate { sequence: u64 },
21}
22
23pub struct ProcessedAdmission<D> {
24    pub key: String,
25    pub delivery: D,
26    pub accepted: bool,
27}
28
29#[derive(Clone, Copy, Debug, Eq, PartialEq)]
30pub enum ControlSignal {
31    Stop,
32    Deadline,
33}
34
35pub enum TurnExit<D> {
36    Complete {
37        answer: Option<String>,
38        processed: Vec<ProcessedAdmission<D>>,
39    },
40    Interrupted {
41        signal: ControlSignal,
42        processed: Vec<ProcessedAdmission<D>>,
43    },
44}
45
46pub struct DriveFailure<D> {
47    pub error: anyhow::Error,
48    pub processed: Vec<ProcessedAdmission<D>>,
49}
50
51impl<D> std::fmt::Debug for DriveFailure<D> {
52    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
53        write!(
54            f,
55            "DriveFailure {{ error: {:?}, processed: {} admissions }}",
56            self.error,
57            self.processed.len()
58        )
59    }
60}
61
62impl<D> std::fmt::Display for DriveFailure<D> {
63    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
64        std::fmt::Display::fmt(&self.error, f)
65    }
66}
67
68impl<D> std::error::Error for DriveFailure<D> {
69    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
70        Some(self.error.as_ref())
71    }
72}
73
74#[derive(Clone, Copy, Eq, PartialEq)]
75enum ClaimState {
76    Queued,
77    Drained,
78}
79
80#[derive(Clone, Copy)]
81struct Claim {
82    sequence: u64,
83    state: ClaimState,
84}
85
86struct SequencedAdmission<D> {
87    sequence: u64,
88    item: QueuedAdmission<D>,
89}
90
91struct MailboxState<D> {
92    last_sequence: u64,
93    queue: VecDeque<SequencedAdmission<D>>,
94    claims: HashMap<String, Claim>,
95}
96
97struct Shared<D> {
98    state: Mutex<MailboxState<D>>,
99    notify: Notify,
100}
101
102pub struct Mailbox<D> {
103    shared: Arc<Shared<D>>,
104}
105
106pub struct MailboxSender<D> {
107    shared: Arc<Shared<D>>,
108}
109
110impl<D> Clone for MailboxSender<D> {
111    fn clone(&self) -> Self {
112        Self {
113            shared: Arc::clone(&self.shared),
114        }
115    }
116}
117
118impl<D> Mailbox<D> {
119    pub fn new() -> Self {
120        Self {
121            shared: Arc::new(Shared {
122                state: Mutex::new(MailboxState {
123                    last_sequence: 0,
124                    queue: VecDeque::new(),
125                    claims: HashMap::new(),
126                }),
127                notify: Notify::new(),
128            }),
129        }
130    }
131
132    fn lock(&self) -> MutexGuard<'_, MailboxState<D>> {
133        self.shared
134            .state
135            .lock()
136            .unwrap_or_else(|error| error.into_inner())
137    }
138
139    pub fn sender(&self) -> MailboxSender<D> {
140        MailboxSender {
141            shared: Arc::clone(&self.shared),
142        }
143    }
144
145    pub fn is_empty(&self) -> bool {
146        self.lock().queue.is_empty()
147    }
148
149    pub async fn notified(&self) {
150        self.shared.notify.notified().await;
151    }
152
153    fn watermark(&self) -> u64 {
154        self.lock().last_sequence
155    }
156
157    fn drain_through(&self, watermark: u64) -> Vec<SequencedAdmission<D>> {
158        let mut state = self.lock();
159        let mut drained = Vec::new();
160        while state
161            .queue
162            .front()
163            .is_some_and(|item| item.sequence <= watermark)
164        {
165            let item = state.queue.pop_front().expect("front was present");
166            let claim = state
167                .claims
168                .get_mut(&item.item.key)
169                .expect("queued item had no claim");
170            assert!(
171                claim.sequence == item.sequence && claim.state == ClaimState::Queued,
172                "queued item had an invalid claim"
173            );
174            claim.state = ClaimState::Drained;
175            drained.push(item);
176        }
177        drained
178    }
179
180    fn restore_front(&self, items: Vec<SequencedAdmission<D>>) {
181        {
182            let mut state = self.lock();
183            let mut previous = None;
184            for item in &items {
185                assert!(
186                    previous.is_none_or(|sequence| sequence < item.sequence),
187                    "privately drained batch was out of order"
188                );
189                previous = Some(item.sequence);
190                assert!(
191                    state.claims.get(&item.item.key).is_some_and(|claim| {
192                        claim.sequence == item.sequence && claim.state == ClaimState::Drained
193                    }),
194                    "privately drained item had an invalid claim"
195                );
196            }
197            if let (Some(last), Some(front)) = (items.last(), state.queue.front()) {
198                assert!(
199                    last.sequence < front.sequence,
200                    "privately drained batch did not precede later arrivals"
201                );
202            }
203            for item in &items {
204                state
205                    .claims
206                    .get_mut(&item.item.key)
207                    .expect("asserted claim was absent")
208                    .state = ClaimState::Queued;
209            }
210            for item in items.into_iter().rev() {
211                state.queue.push_front(item);
212            }
213        }
214        self.shared.notify.notify_one();
215    }
216
217    fn acknowledge(&self, sequence: u64, key: &str) {
218        let mut state = self.lock();
219        assert!(
220            state.claims.get(key).is_some_and(|claim| {
221                claim.sequence == sequence && claim.state == ClaimState::Drained
222            }),
223            "privately drained item had an invalid claim"
224        );
225        state.claims.remove(key).expect("asserted claim was absent");
226    }
227}
228
229impl<D> Default for Mailbox<D> {
230    fn default() -> Self {
231        Self::new()
232    }
233}
234
235impl<D> MailboxSender<D> {
236    pub fn push(&self, item: QueuedAdmission<D>) -> PushResult {
237        let result = {
238            let mut state = self
239                .shared
240                .state
241                .lock()
242                .unwrap_or_else(|error| error.into_inner());
243            if let Some(claim) = state.claims.get(&item.key) {
244                return PushResult::Duplicate {
245                    sequence: claim.sequence,
246                };
247            }
248            let sequence = state
249                .last_sequence
250                .checked_add(1)
251                .expect("mailbox sequence exhausted");
252            state.last_sequence = sequence;
253            state.claims.insert(
254                item.key.clone(),
255                Claim {
256                    sequence,
257                    state: ClaimState::Queued,
258                },
259            );
260            state.queue.push_back(SequencedAdmission { sequence, item });
261            PushResult::Queued { sequence }
262        };
263        self.shared.notify.notify_one();
264        result
265    }
266}
267
268fn clone_admission(admission: &PendingTurnAdmission) -> PendingTurnAdmission {
269    match admission {
270        PendingTurnAdmission::User { text, metadata } => PendingTurnAdmission::User {
271            text: text.clone(),
272            metadata: metadata.clone(),
273        },
274        PendingTurnAdmission::Source {
275            kennedy,
276            text,
277            metadata,
278        } => PendingTurnAdmission::Source {
279            kennedy: *kennedy,
280            text: text.clone(),
281            metadata: metadata.clone(),
282        },
283    }
284}
285
286fn failed<D>(error: anyhow::Error, processed: Vec<ProcessedAdmission<D>>) -> DriveFailure<D> {
287    DriveFailure { error, processed }
288}
289
290pub async fn drive_pending_turn<D, C, F, S, K>(
291    session: &mut Session,
292    operation_id: Uuid,
293    turn_deadline: Option<TurnDeadline>,
294    mailbox: &mut Mailbox<D>,
295    control: S,
296    mut checkpoint: C,
297    mut cancel: K,
298) -> Result<TurnExit<D>, DriveFailure<D>>
299where
300    D: Send,
301    C: FnMut(Value) -> F + Send,
302    F: Future<Output = anyhow::Result<()>> + Send,
303    S: Future<Output = ControlSignal> + Send,
304    K: FnMut(ControlSignal),
305{
306    let mut processed = Vec::new();
307    let turn = match session.begin_pending_turn(operation_id, turn_deadline) {
308        Ok(Some(turn)) => turn,
309        Ok(None) => {
310            return Ok(TurnExit::Complete {
311                answer: None,
312                processed,
313            });
314        }
315        Err(error) => return Err(failed(error, processed)),
316    };
317    tokio::pin!(control);
318    let mut boundary = match session.advance_pending_turn(turn, &mut checkpoint).await {
319        Ok(boundary) => boundary,
320        Err(error) => return Err(failed(error, processed)),
321    };
322    loop {
323        match boundary {
324            TurnBoundary::Complete(answer) => {
325                return Ok(TurnExit::Complete { answer, processed });
326            }
327            TurnBoundary::Yield(mut turn) => {
328                let watermark = mailbox.watermark();
329                let mut remaining = mailbox.drain_through(watermark).into_iter();
330                while let Some(item) = remaining.next() {
331                    let admission = clone_admission(&item.item.admission);
332                    match session
333                        .admit_pending_turn(
334                            &mut turn,
335                            admission,
336                            &item.item.recorded_at,
337                            &mut checkpoint,
338                        )
339                        .await
340                    {
341                        Ok(accepted) => {
342                            let sequence = item.sequence;
343                            let QueuedAdmission { key, delivery, .. } = item.item;
344                            mailbox.acknowledge(sequence, &key);
345                            processed.push(ProcessedAdmission {
346                                key,
347                                delivery,
348                                accepted,
349                            });
350                        }
351                        Err(error) => {
352                            let mut restore = vec![item];
353                            restore.extend(remaining);
354                            mailbox.restore_front(restore);
355                            return Err(failed(error, processed));
356                        }
357                    }
358                }
359                boundary = match session.advance_pending_turn(turn, &mut checkpoint).await {
360                    Ok(boundary) => boundary,
361                    Err(error) => return Err(failed(error, processed)),
362                };
363            }
364            TurnBoundary::Await(pending) => {
365                let mut waiter = tokio::spawn(pending.wait());
366                tokio::select! {
367                    biased;
368                    signal = &mut control => {
369                        cancel(signal);
370                        waiter.abort();
371                        match waiter.await {
372                            Err(error) if error.is_cancelled() => {}
373                            Err(error) => return Err(failed(error.into(), processed)),
374                            Ok(_) => {}
375                        }
376                        return Ok(TurnExit::Interrupted { signal, processed });
377                    }
378                    joined = &mut waiter => {
379                        let wake = match joined {
380                            Ok(wake) => wake,
381                            Err(error) => return Err(failed(error.into(), processed)),
382                        };
383                        boundary = match session
384                            .apply_inference_wake(wake, &mut checkpoint)
385                            .await
386                        {
387                            Ok(boundary) => boundary,
388                            Err(error) => return Err(failed(error, processed)),
389                        };
390                    }
391                }
392            }
393        }
394    }
395}
396
397#[cfg(test)]
398mod tests;