kcode_kennedy_stepped_turn_runtime/
lib.rs1pub 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;