kcode_k1_chat_thread_session_view/
lib.rs1#![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}