Skip to main content

kcode_k1_chat_thread_session_view/
lib.rs

1#![forbid(unsafe_code)]
2
3use kcode_k1_chat_thread_actor_channel::{
4    ActorError, ActorStatus, ProviderInput, Reply, Snapshot, SnapshotWithEvents,
5};
6use kcode_k1_chat_thread_durable_state::{DurableThread, Status};
7use tokio::sync::mpsc;
8
9pub struct SessionView {
10    event_receiver: Option<mpsc::UnboundedReceiver<Reply<SnapshotWithEvents>>>,
11    model_input: Option<ProviderInput>,
12    waiters: Vec<Reply<Snapshot>>,
13}
14
15impl SessionView {
16    pub fn new(event_receiver: mpsc::UnboundedReceiver<Reply<SnapshotWithEvents>>) -> Self {
17        Self {
18            event_receiver: Some(event_receiver),
19            model_input: None,
20            waiters: Vec::new(),
21        }
22    }
23
24    pub fn set_model_input(&mut self, input: ProviderInput) {
25        self.model_input = Some(input);
26    }
27
28    pub fn snapshot(&self, durable: &DurableThread, force_running: bool) -> Snapshot {
29        Snapshot {
30            boxes: durable.boxes().to_vec(),
31            status: if force_running {
32                ActorStatus::Running
33            } else {
34                durable.status()
35            },
36            model_input: self.model_input.clone(),
37        }
38    }
39
40    pub fn wait(&mut self, durable: &DurableThread, force_running: bool, reply: Reply<Snapshot>) {
41        if running(durable, force_running) {
42            self.waiters.push(reply);
43        } else {
44            let _ = reply.send(Ok(self.snapshot(durable, force_running)));
45        }
46    }
47
48    pub fn wake(&mut self, durable: &DurableThread, force_running: bool) {
49        if running(durable, force_running) {
50            return;
51        }
52        let snapshot = self.snapshot(durable, force_running);
53        for reply in std::mem::take(&mut self.waiters) {
54            let _ = reply.send(Ok(snapshot.clone()));
55        }
56    }
57
58    pub async fn receive_event_query(&mut self) -> Option<Reply<SnapshotWithEvents>> {
59        if let Some(receiver) = self.event_receiver.as_mut() {
60            if let Some(reply) = receiver.recv().await {
61                return Some(reply);
62            }
63            self.event_receiver = None;
64        }
65        std::future::pending().await
66    }
67
68    pub fn answer_event_query(
69        &self,
70        durable: &DurableThread,
71        force_running: bool,
72        reply: Reply<SnapshotWithEvents>,
73    ) {
74        let _ = reply.send(Ok(SnapshotWithEvents {
75            snapshot: self.snapshot(durable, force_running),
76            events: durable.events(),
77        }));
78    }
79
80    pub fn close(&mut self) {
81        for reply in std::mem::take(&mut self.waiters) {
82            let _ = reply.send(Err(ActorError::Closed));
83        }
84    }
85}
86
87fn running(durable: &DurableThread, force_running: bool) -> bool {
88    force_running || matches!(durable.status(), Status::Running)
89}
90
91#[cfg(test)]
92mod tests {
93    use super::*;
94    use kcode_k1_chat_thread_actor_channel::{ActorError, ProviderInputKind, channel_with_events};
95    use tokio::sync::oneshot;
96
97    #[test]
98    fn model_input_retains_the_latest_value() {
99        let (_sender, receiver) = mpsc::unbounded_channel();
100        let mut view = SessionView::new(receiver);
101        let first = ProviderInput {
102            kind: ProviderInputKind::Turn,
103            text: "first".into(),
104        };
105        let latest = ProviderInput {
106            kind: ProviderInputKind::MailboxFlush,
107            text: "latest".into(),
108        };
109        view.set_model_input(first);
110        view.set_model_input(latest.clone());
111        assert_eq!(view.model_input, Some(latest));
112    }
113
114    #[tokio::test]
115    async fn event_queries_arrive_on_the_dedicated_receiver() {
116        let (handle, _sender, _receiver, event_receiver) = channel_with_events();
117        let mut view = SessionView::new(event_receiver);
118        let query = tokio::spawn(async move { handle.snapshot_with_events().await });
119        let reply = view.receive_event_query().await.unwrap();
120        assert!(reply.send(Err(ActorError::NotStalled)).is_ok());
121        assert_eq!(query.await.unwrap(), Err(ActorError::NotStalled));
122    }
123
124    #[tokio::test]
125    async fn closed_event_receiver_remains_pending() {
126        let (handle, _sender, _receiver, event_receiver) = channel_with_events();
127        drop(handle);
128        let mut view = SessionView::new(event_receiver);
129        {
130            let first = view.receive_event_query();
131            tokio::pin!(first);
132            tokio::select! {
133                biased;
134                _ = &mut first => panic!("closed receiver returned"),
135                _ = async {} => {}
136            }
137        }
138        assert!(view.event_receiver.is_none());
139
140        let again = view.receive_event_query();
141        tokio::pin!(again);
142        tokio::select! {
143            biased;
144            _ = &mut again => panic!("retained closure returned"),
145            _ = async {} => {}
146        }
147    }
148
149    #[tokio::test]
150    async fn close_rejects_inserted_waiters() {
151        let (_sender, receiver) = mpsc::unbounded_channel();
152        let mut view = SessionView::new(receiver);
153        let (reply, answer) = oneshot::channel();
154        view.waiters.push(reply);
155        view.close();
156        assert_eq!(answer.await, Ok(Err(ActorError::Closed)));
157    }
158}