kcode-kennedy-stepped-turn-runtime 0.1.1

Watermark-bounded mailbox sequencing and stepped-turn session driving
Documentation
pub use kcode_kennedy_sessions::PendingTurnAdmission;
use kcode_kennedy_sessions::{Session, TurnBoundary, TurnDeadline};
use serde_json::Value;
use std::collections::{HashMap, VecDeque};
use std::future::Future;
use std::sync::{Arc, Mutex, MutexGuard};
use tokio::sync::Notify;
use uuid::Uuid;

pub struct QueuedAdmission<D> {
    pub key: String,
    pub recorded_at: String,
    pub admission: PendingTurnAdmission,
    pub delivery: D,
}

#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum PushResult {
    Queued { sequence: u64 },
    Duplicate { sequence: u64 },
}

pub struct ProcessedAdmission<D> {
    pub key: String,
    pub delivery: D,
    pub accepted: bool,
}

#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum ControlSignal {
    Stop,
    Deadline,
}

pub enum TurnExit<D> {
    Complete {
        answer: Option<String>,
        processed: Vec<ProcessedAdmission<D>>,
    },
    Interrupted {
        signal: ControlSignal,
        processed: Vec<ProcessedAdmission<D>>,
    },
}

pub struct DriveFailure<D> {
    pub error: anyhow::Error,
    pub processed: Vec<ProcessedAdmission<D>>,
}

impl<D> std::fmt::Debug for DriveFailure<D> {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        write!(
            f,
            "DriveFailure {{ error: {:?}, processed: {} admissions }}",
            self.error,
            self.processed.len()
        )
    }
}

impl<D> std::fmt::Display for DriveFailure<D> {
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
        std::fmt::Display::fmt(&self.error, f)
    }
}

impl<D> std::error::Error for DriveFailure<D> {
    fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
        Some(self.error.as_ref())
    }
}

#[derive(Clone, Copy, Eq, PartialEq)]
enum ClaimState {
    Queued,
    Drained,
}

#[derive(Clone, Copy)]
struct Claim {
    sequence: u64,
    state: ClaimState,
}

struct SequencedAdmission<D> {
    sequence: u64,
    item: QueuedAdmission<D>,
}

struct MailboxState<D> {
    last_sequence: u64,
    queue: VecDeque<SequencedAdmission<D>>,
    claims: HashMap<String, Claim>,
}

struct Shared<D> {
    state: Mutex<MailboxState<D>>,
    notify: Notify,
}

pub struct Mailbox<D> {
    shared: Arc<Shared<D>>,
}

pub struct MailboxSender<D> {
    shared: Arc<Shared<D>>,
}

impl<D> Clone for MailboxSender<D> {
    fn clone(&self) -> Self {
        Self {
            shared: Arc::clone(&self.shared),
        }
    }
}

impl<D> Mailbox<D> {
    pub fn new() -> Self {
        Self {
            shared: Arc::new(Shared {
                state: Mutex::new(MailboxState {
                    last_sequence: 0,
                    queue: VecDeque::new(),
                    claims: HashMap::new(),
                }),
                notify: Notify::new(),
            }),
        }
    }

    fn lock(&self) -> MutexGuard<'_, MailboxState<D>> {
        self.shared
            .state
            .lock()
            .unwrap_or_else(|error| error.into_inner())
    }

    pub fn sender(&self) -> MailboxSender<D> {
        MailboxSender {
            shared: Arc::clone(&self.shared),
        }
    }

    pub fn is_empty(&self) -> bool {
        self.lock().queue.is_empty()
    }

    pub async fn notified(&self) {
        self.shared.notify.notified().await;
    }

    fn watermark(&self) -> u64 {
        self.lock().last_sequence
    }

    fn drain_through(&self, watermark: u64) -> Vec<SequencedAdmission<D>> {
        let mut state = self.lock();
        let mut drained = Vec::new();
        while state
            .queue
            .front()
            .is_some_and(|item| item.sequence <= watermark)
        {
            let item = state.queue.pop_front().expect("front was present");
            let claim = state
                .claims
                .get_mut(&item.item.key)
                .expect("queued item had no claim");
            assert!(
                claim.sequence == item.sequence && claim.state == ClaimState::Queued,
                "queued item had an invalid claim"
            );
            claim.state = ClaimState::Drained;
            drained.push(item);
        }
        drained
    }

    fn restore_front(&self, items: Vec<SequencedAdmission<D>>) {
        {
            let mut state = self.lock();
            let mut previous = None;
            for item in &items {
                assert!(
                    previous.is_none_or(|sequence| sequence < item.sequence),
                    "privately drained batch was out of order"
                );
                previous = Some(item.sequence);
                assert!(
                    state.claims.get(&item.item.key).is_some_and(|claim| {
                        claim.sequence == item.sequence && claim.state == ClaimState::Drained
                    }),
                    "privately drained item had an invalid claim"
                );
            }
            if let (Some(last), Some(front)) = (items.last(), state.queue.front()) {
                assert!(
                    last.sequence < front.sequence,
                    "privately drained batch did not precede later arrivals"
                );
            }
            for item in &items {
                state
                    .claims
                    .get_mut(&item.item.key)
                    .expect("asserted claim was absent")
                    .state = ClaimState::Queued;
            }
            for item in items.into_iter().rev() {
                state.queue.push_front(item);
            }
        }
        self.shared.notify.notify_one();
    }

    fn acknowledge(&self, sequence: u64, key: &str) {
        let mut state = self.lock();
        assert!(
            state.claims.get(key).is_some_and(|claim| {
                claim.sequence == sequence && claim.state == ClaimState::Drained
            }),
            "privately drained item had an invalid claim"
        );
        state.claims.remove(key).expect("asserted claim was absent");
    }
}

impl<D> Default for Mailbox<D> {
    fn default() -> Self {
        Self::new()
    }
}

impl<D> MailboxSender<D> {
    pub fn push(&self, item: QueuedAdmission<D>) -> PushResult {
        let result = {
            let mut state = self
                .shared
                .state
                .lock()
                .unwrap_or_else(|error| error.into_inner());
            if let Some(claim) = state.claims.get(&item.key) {
                return PushResult::Duplicate {
                    sequence: claim.sequence,
                };
            }
            let sequence = state
                .last_sequence
                .checked_add(1)
                .expect("mailbox sequence exhausted");
            state.last_sequence = sequence;
            state.claims.insert(
                item.key.clone(),
                Claim {
                    sequence,
                    state: ClaimState::Queued,
                },
            );
            state.queue.push_back(SequencedAdmission { sequence, item });
            PushResult::Queued { sequence }
        };
        self.shared.notify.notify_one();
        result
    }
}

fn clone_admission(admission: &PendingTurnAdmission) -> PendingTurnAdmission {
    match admission {
        PendingTurnAdmission::User { text, metadata } => PendingTurnAdmission::User {
            text: text.clone(),
            metadata: metadata.clone(),
        },
        PendingTurnAdmission::Source {
            kennedy,
            text,
            metadata,
        } => PendingTurnAdmission::Source {
            kennedy: *kennedy,
            text: text.clone(),
            metadata: metadata.clone(),
        },
    }
}

fn failed<D>(error: anyhow::Error, processed: Vec<ProcessedAdmission<D>>) -> DriveFailure<D> {
    DriveFailure { error, processed }
}

pub async fn drive_pending_turn<D, C, F, S, K>(
    session: &mut Session,
    operation_id: Uuid,
    turn_deadline: Option<TurnDeadline>,
    mailbox: &mut Mailbox<D>,
    control: S,
    mut checkpoint: C,
    mut cancel: K,
) -> Result<TurnExit<D>, DriveFailure<D>>
where
    D: Send,
    C: FnMut(Value) -> F + Send,
    F: Future<Output = anyhow::Result<()>> + Send,
    S: Future<Output = ControlSignal> + Send,
    K: FnMut(ControlSignal),
{
    let mut processed = Vec::new();
    let turn = match session.begin_pending_turn(operation_id, turn_deadline) {
        Ok(Some(turn)) => turn,
        Ok(None) => {
            return Ok(TurnExit::Complete {
                answer: None,
                processed,
            });
        }
        Err(error) => return Err(failed(error, processed)),
    };
    tokio::pin!(control);
    let mut boundary = match session.advance_pending_turn(turn, &mut checkpoint).await {
        Ok(boundary) => boundary,
        Err(error) => return Err(failed(error, processed)),
    };
    loop {
        match boundary {
            TurnBoundary::Complete(answer) => {
                return Ok(TurnExit::Complete { answer, processed });
            }
            TurnBoundary::Yield(mut turn) => {
                let watermark = mailbox.watermark();
                let mut remaining = mailbox.drain_through(watermark).into_iter();
                while let Some(item) = remaining.next() {
                    let admission = clone_admission(&item.item.admission);
                    match session
                        .admit_pending_turn(
                            &mut turn,
                            admission,
                            &item.item.recorded_at,
                            &mut checkpoint,
                        )
                        .await
                    {
                        Ok(accepted) => {
                            let sequence = item.sequence;
                            let QueuedAdmission { key, delivery, .. } = item.item;
                            mailbox.acknowledge(sequence, &key);
                            processed.push(ProcessedAdmission {
                                key,
                                delivery,
                                accepted,
                            });
                        }
                        Err(error) => {
                            let mut restore = vec![item];
                            restore.extend(remaining);
                            mailbox.restore_front(restore);
                            return Err(failed(error, processed));
                        }
                    }
                }
                boundary = match session.advance_pending_turn(turn, &mut checkpoint).await {
                    Ok(boundary) => boundary,
                    Err(error) => return Err(failed(error, processed)),
                };
            }
            TurnBoundary::Await(pending) => {
                let mut waiter = tokio::spawn(pending.wait());
                tokio::select! {
                    biased;
                    signal = &mut control => {
                        cancel(signal);
                        waiter.abort();
                        match waiter.await {
                            Err(error) if error.is_cancelled() => {}
                            Err(error) => return Err(failed(error.into(), processed)),
                            Ok(_) => {}
                        }
                        return Ok(TurnExit::Interrupted { signal, processed });
                    }
                    joined = &mut waiter => {
                        let wake = match joined {
                            Ok(wake) => wake,
                            Err(error) => return Err(failed(error.into(), processed)),
                        };
                        boundary = match session
                            .apply_inference_wake(wake, &mut checkpoint)
                            .await
                        {
                            Ok(boundary) => boundary,
                            Err(error) => return Err(failed(error, processed)),
                        };
                    }
                }
            }
        }
    }
}

#[cfg(test)]
mod tests;