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;