kcode-k1-chat-thread-session-view 0.2.1

Process-local read projection for one K1 chat-thread session actor
Documentation
#![forbid(unsafe_code)]

use kcode_k1_chat_thread_actor_channel::{
    ActorError, ActorStatus, ProviderInput, Reply, Snapshot, SnapshotWithEvents,
};
use kcode_k1_chat_thread_durable_state::{DurableThread, Status};
use tokio::sync::mpsc;

pub struct SessionView {
    event_receiver: Option<mpsc::UnboundedReceiver<Reply<SnapshotWithEvents>>>,
    model_input: Option<ProviderInput>,
    waiters: Vec<Reply<Snapshot>>,
}

impl SessionView {
    pub fn new(event_receiver: mpsc::UnboundedReceiver<Reply<SnapshotWithEvents>>) -> Self {
        Self {
            event_receiver: Some(event_receiver),
            model_input: None,
            waiters: Vec::new(),
        }
    }

    pub fn set_model_input(&mut self, input: ProviderInput) {
        self.model_input = Some(input);
    }

    pub fn snapshot(&self, durable: &DurableThread, force_running: bool) -> Snapshot {
        Snapshot {
            boxes: durable.boxes().to_vec(),
            status: if force_running {
                ActorStatus::Running
            } else {
                durable.status()
            },
            model_input: self.model_input.clone(),
        }
    }

    pub fn wait(&mut self, durable: &DurableThread, force_running: bool, reply: Reply<Snapshot>) {
        if running(durable, force_running) {
            self.waiters.push(reply);
        } else {
            let _ = reply.send(Ok(self.snapshot(durable, force_running)));
        }
    }

    pub fn wake(&mut self, durable: &DurableThread, force_running: bool) {
        if running(durable, force_running) {
            return;
        }
        let snapshot = self.snapshot(durable, force_running);
        for reply in std::mem::take(&mut self.waiters) {
            let _ = reply.send(Ok(snapshot.clone()));
        }
    }

    pub async fn receive_event_query(&mut self) -> Option<Reply<SnapshotWithEvents>> {
        if let Some(receiver) = self.event_receiver.as_mut() {
            if let Some(reply) = receiver.recv().await {
                return Some(reply);
            }
            self.event_receiver = None;
        }
        std::future::pending().await
    }

    pub fn answer_event_query(
        &self,
        durable: &DurableThread,
        force_running: bool,
        reply: Reply<SnapshotWithEvents>,
    ) {
        let _ = reply.send(Ok(SnapshotWithEvents {
            snapshot: self.snapshot(durable, force_running),
            events: durable.events(),
        }));
    }

    pub fn close(&mut self) {
        for reply in std::mem::take(&mut self.waiters) {
            let _ = reply.send(Err(ActorError::Closed));
        }
    }
}

fn running(durable: &DurableThread, force_running: bool) -> bool {
    force_running || matches!(durable.status(), Status::Running)
}

#[cfg(test)]
mod tests {
    use super::*;
    use kcode_k1_chat_thread_actor_channel::{ActorError, ProviderInputKind, channel_with_events};
    use tokio::sync::oneshot;

    #[test]
    fn model_input_retains_the_latest_value() {
        let (_sender, receiver) = mpsc::unbounded_channel();
        let mut view = SessionView::new(receiver);
        let first = ProviderInput {
            kind: ProviderInputKind::Turn,
            text: "first".into(),
        };
        let latest = ProviderInput {
            kind: ProviderInputKind::MailboxFlush,
            text: "latest".into(),
        };
        view.set_model_input(first);
        view.set_model_input(latest.clone());
        assert_eq!(view.model_input, Some(latest));
    }

    #[tokio::test]
    async fn event_queries_arrive_on_the_dedicated_receiver() {
        let (handle, _sender, _receiver, event_receiver) = channel_with_events();
        let mut view = SessionView::new(event_receiver);
        let query = tokio::spawn(async move { handle.snapshot_with_events().await });
        let reply = view.receive_event_query().await.unwrap();
        assert!(reply.send(Err(ActorError::NotStalled)).is_ok());
        assert_eq!(query.await.unwrap(), Err(ActorError::NotStalled));
    }

    #[tokio::test]
    async fn closed_event_receiver_remains_pending() {
        let (handle, _sender, _receiver, event_receiver) = channel_with_events();
        drop(handle);
        let mut view = SessionView::new(event_receiver);
        {
            let first = view.receive_event_query();
            tokio::pin!(first);
            tokio::select! {
                biased;
                _ = &mut first => panic!("closed receiver returned"),
                _ = async {} => {}
            }
        }
        assert!(view.event_receiver.is_none());

        let again = view.receive_event_query();
        tokio::pin!(again);
        tokio::select! {
            biased;
            _ = &mut again => panic!("retained closure returned"),
            _ = async {} => {}
        }
    }

    #[tokio::test]
    async fn close_rejects_inserted_waiters() {
        let (_sender, receiver) = mpsc::unbounded_channel();
        let mut view = SessionView::new(receiver);
        let (reply, answer) = oneshot::channel();
        view.waiters.push(reply);
        view.close();
        assert_eq!(answer.await, Ok(Err(ActorError::Closed)));
    }
}