1#![forbid(unsafe_code)]
2#![doc = include_str!("../Documentation.md")]
3
4use kcode_k1_access_kmap::K1AccessKmap;
5use kcode_k1_chat_model_usage_session::ModelUsageSession;
6use kcode_k1_chat_persistence::Session;
7use kcode_k1_chat_thread_actor_channel::{
8 ActorError, ActorShim, BoxValue, Handle, Message, ProviderInput, ProviderInputKind, Reply,
9 channel_with_events, new_shim,
10};
11use kcode_k1_chat_thread_durable_state::{
12 AccessContext, AccessPolicy, DurableThread, ProfileId, SetLaunchNodeKtool, TransitionError,
13};
14use kcode_k1_chat_thread_session_code_runtime::SessionCodeRuntime;
15use kcode_k1_chat_thread_session_stage_runtime::{EventContext, StageRuntime};
16use kcode_k1_chat_thread_session_view::SessionView;
17use kcode_k1_codex_adapter::{Adapter, ShimOutput};
18use kcode_k1_codex_websearch::Runner as WebSearchRunner;
19use kcode_k1_rust_code_ktool_service::RustCodeKtoolService;
20use kcode_k1_web_code_ktool_service::K1WebCodeKtoolService;
21use std::sync::Arc;
22use tokio::{sync::mpsc, task::JoinHandle};
23
24type ActorResult = Result<bool, String>;
25type InferenceResult = Result<ShimOutput<BoxValue>, String>;
26type UnitResult = Result<(), String>;
27
28pub fn open(
29 adapter: Adapter,
30 key: impl Into<String>,
31 session: Session,
32 kmap: Arc<K1AccessKmap>,
33 web_search: WebSearchRunner,
34) -> Result<Handle, String> {
35 let durable = DurableThread::recover(session, kmap)?;
36 spawn_actor(adapter, key, durable, None, None, web_search)
37}
38
39pub fn open_with_social(
40 adapter: Adapter,
41 key: impl Into<String>,
42 session: Session,
43 kmap: Arc<K1AccessKmap>,
44 social: kcode_k1_ktool_social::SocialKtools,
45 web_search: WebSearchRunner,
46) -> Result<Handle, String> {
47 let durable = DurableThread::recover_with_social(session, kmap, social)?;
48 spawn_actor(adapter, key, durable, None, None, web_search)
49}
50
51pub fn open_with_social_and_set_launch_node(
52 adapter: Adapter,
53 key: impl Into<String>,
54 session: Session,
55 kmap: Arc<K1AccessKmap>,
56 social: kcode_k1_ktool_social::SocialKtools,
57 set_launch_node: SetLaunchNodeKtool,
58 web_search: WebSearchRunner,
59) -> Result<Handle, String> {
60 let durable = DurableThread::recover_with_social_and_set_launch_node(
61 session,
62 kmap,
63 social,
64 set_launch_node,
65 )?;
66 spawn_actor(adapter, key, durable, None, None, web_search)
67}
68
69#[allow(clippy::too_many_arguments)]
70pub fn open_with_social_and_set_launch_node_and_rust_code(
71 adapter: Adapter,
72 key: impl Into<String>,
73 session: Session,
74 kmap: Arc<K1AccessKmap>,
75 social: kcode_k1_ktool_social::SocialKtools,
76 set_launch_node: SetLaunchNodeKtool,
77 rust_code: Arc<RustCodeKtoolService>,
78 web_search: WebSearchRunner,
79) -> Result<Handle, String> {
80 let durable = DurableThread::recover_with_social_and_set_launch_node(
81 session,
82 kmap,
83 social,
84 set_launch_node,
85 )?;
86 spawn_actor(adapter, key, durable, Some(rust_code), None, web_search)
87}
88
89#[allow(clippy::too_many_arguments)]
90pub fn open_with_social_and_set_launch_node_and_rust_code_and_web_code(
91 adapter: Adapter,
92 key: impl Into<String>,
93 session: Session,
94 kmap: Arc<K1AccessKmap>,
95 social: kcode_k1_ktool_social::SocialKtools,
96 set_launch_node: SetLaunchNodeKtool,
97 rust_code: Arc<RustCodeKtoolService>,
98 web_code: K1WebCodeKtoolService,
99 web_search: WebSearchRunner,
100) -> Result<Handle, String> {
101 let durable = DurableThread::recover_with_social_and_set_launch_node(
102 session,
103 kmap,
104 social,
105 set_launch_node,
106 )?;
107 spawn_actor(
108 adapter,
109 key,
110 durable,
111 Some(rust_code),
112 Some(web_code),
113 web_search,
114 )
115}
116
117fn spawn_actor(
118 adapter: Adapter,
119 key: impl Into<String>,
120 durable: DurableThread,
121 rust_code: Option<Arc<RustCodeKtoolService>>,
122 web_code: Option<K1WebCodeKtoolService>,
123 web_search: WebSearchRunner,
124) -> Result<Handle, String> {
125 let code = SessionCodeRuntime::recover(&durable, rust_code, web_code)?;
126 let stage = StageRuntime::new(code, web_search);
127 let (handle, sender, receiver, event_receiver) = channel_with_events();
128 let base_key = key.into();
129 let actor = Actor {
130 durable,
131 shim: Some(new_shim(adapter.clone(), base_key.clone(), sender.clone())),
132 adapter,
133 active_key: base_key.clone(),
134 base_key,
135 generation: 0,
136 stage,
137 view: SessionView::new(event_receiver),
138 sender,
139 receiver,
140 usage: ModelUsageSession::default(),
141 inference: None,
142 job: None,
143 };
144 tokio::spawn(actor.run());
145 Ok(handle)
146}
147
148struct Actor {
149 durable: DurableThread,
150 shim: Option<ActorShim>,
151 adapter: Adapter,
152 base_key: String,
153 active_key: String,
154 generation: u64,
155 stage: StageRuntime,
156 view: SessionView,
157 sender: mpsc::UnboundedSender<Message>,
158 receiver: mpsc::UnboundedReceiver<Message>,
159 usage: ModelUsageSession,
160 inference: Option<JoinHandle<()>>,
161 job: Option<u64>,
162}
163
164impl Actor {
165 async fn run(mut self) {
166 loop {
167 if let Err(error) = self.drive().await {
168 self.fatal(error);
169 break;
170 }
171 let searching = self.stage.search_active();
172 let can_receive_code = self.stage.can_receive_code();
173 self.view.wake(&self.durable, searching);
174 let stop = tokio::select! {
175 message = self.receiver.recv() => match message {
176 Some(message) => match self.handle(message).await {
177 Ok(stop) => stop,
178 Err(error) => self.fatal(error),
179 },
180 None => true,
181 },
182 result = self.stage.receive_code(), if can_receive_code => match result {
183 Ok(()) => match self.handle_code_completion() {
184 Ok(stop) => stop,
185 Err(error) => self.fatal(error),
186 },
187 Err(error) => self.fatal(error),
188 },
189 reply = self.view.receive_event_query() => {
190 if let Some(reply) = reply {
191 self.view.answer_event_query(&self.durable, searching, reply);
192 }
193 false
194 }
195 usage = self.usage.receive(&mut self.durable) => usage.is_err(),
196 };
197 if stop {
198 break;
199 }
200 }
201 self.stage.abort();
202 let inference_active = self.inference.is_some();
203 let _ = self
204 .usage
205 .shutdown(&mut self.durable, inference_active)
206 .await;
207 self.stage.shutdown();
208 if let Some(task) = self.inference.take() {
209 task.abort();
210 }
211 self.view.close();
212 }
213
214 async fn handle(&mut self, message: Message) -> ActorResult {
215 match message {
216 Message::Accept((box_type, contents, hidden_type, hidden_contents), reply) => {
217 Ok(reply_transition(
218 reply,
219 self.durable.accept_external_box(
220 box_type,
221 contents,
222 hidden_type,
223 hidden_contents,
224 ),
225 ))
226 }
227 Message::AcceptUser(context, profile_id, policy, contents, reply) => {
228 let result = self.stage.accept_user(
229 &mut self.durable,
230 context,
231 profile_id,
232 policy,
233 contents,
234 );
235 Ok(reply_transition(reply, result))
236 }
237 Message::Return(id, result, reply) => {
238 let result = self
239 .durable
240 .accept_return(id, result)
241 .map_err(|_| ActorError::Closed);
242 let stop = result.is_err();
243 Ok(answer(reply, result, stop))
244 }
245 Message::Stage(text, boxes, reply) => {
246 let mut event = EventContext {
247 durable: &mut self.durable,
248 adapter: &self.adapter,
249 active_key: &self.active_key,
250 view: &mut self.view,
251 sender: &self.sender,
252 job: self.job,
253 };
254 self.stage.handle_stage(&mut event, text, boxes, reply)
255 }
256 Message::WebSearchCompleted(epoch, id, result) => {
257 let mut event = EventContext {
258 durable: &mut self.durable,
259 adapter: &self.adapter,
260 active_key: &self.active_key,
261 view: &mut self.view,
262 sender: &self.sender,
263 job: self.job,
264 };
265 self.stage
266 .handle_web_search_completion(&mut event, epoch, id, result)
267 }
268 Message::MailboxFlushCompleted(job, prepared, result) => {
269 let mut event = EventContext {
270 durable: &mut self.durable,
271 adapter: &self.adapter,
272 active_key: &self.active_key,
273 view: &mut self.view,
274 sender: &self.sender,
275 job: self.job,
276 };
277 self.stage
278 .handle_mailbox_flush_completion(&mut event, job, prepared, result)
279 }
280 Message::Inferred(job, shim, result) => self.finish(job, shim, result).await,
281 Message::Snapshot(reply) => {
282 let snapshot = self
283 .view
284 .snapshot(&self.durable, self.stage.search_active());
285 Ok(answer(reply, Ok(snapshot), false))
286 }
287 Message::Wait(reply) => {
288 self.view
289 .wait(&self.durable, self.stage.search_active(), reply);
290 Ok(false)
291 }
292 Message::Restart(context, profile_id, policy, reply) => {
293 let result = self.restart(context, profile_id, policy);
294 if let Err(error) = result {
295 return Ok(answer(reply, Err(error), false));
296 }
297 if self.usage.restart(&mut self.durable).await.is_err() {
298 Ok(answer(reply, Err(ActorError::Closed), true))
299 } else {
300 Ok(answer(reply, Ok(()), false))
301 }
302 }
303 Message::Abandon => Ok(true),
304 }
305 }
306
307 fn handle_code_completion(&mut self) -> ActorResult {
308 let mut event = EventContext {
309 durable: &mut self.durable,
310 adapter: &self.adapter,
311 active_key: &self.active_key,
312 view: &mut self.view,
313 sender: &self.sender,
314 job: self.job,
315 };
316 self.stage.handle_code_completion(&mut event)
317 }
318
319 fn fatal(&mut self, error: String) -> bool {
320 self.stage.fail(error)
321 }
322
323 async fn drive(&mut self) -> UnitResult {
324 if self.inference.is_some() || self.shim.is_none() {
325 return Ok(());
326 }
327 let Some((job, input)) = self.durable.begin_input()? else {
328 return Ok(());
329 };
330 if !self.usage.is_subscribed() {
331 self.usage
332 .subscribe(&self.adapter, self.active_key.clone())
333 .await
334 .map_err(|_| "model-usage subscription failed".to_owned())?;
335 }
336 self.view.set_model_input(ProviderInput {
337 kind: ProviderInputKind::Turn,
338 text: input.clone(),
339 });
340 let mut shim = self.shim.take().expect("shim was checked");
341 self.job = Some(job);
342 let sender = self.sender.clone();
343 self.inference = Some(tokio::spawn(async move {
344 let result = shim.infer(input).await.map_err(|error| error.to_string());
345 let _ = sender.send(Message::Inferred(job, shim, result));
346 }));
347 Ok(())
348 }
349
350 async fn finish(&mut self, job: u64, shim: ActorShim, result: InferenceResult) -> ActorResult {
351 if self.inference.take().is_none() || self.stage.mailbox_pending() || self.job != Some(job)
352 {
353 return Err("stale Codex inference completion".to_owned());
354 }
355 if self.stage.code_pending() {
356 return Err("Codex inference completed while Rust code work remains".to_owned());
357 }
358 self.job = None;
359 let key = self.active_key.clone();
360 let (error, resume, terminal_id) = match result {
361 Err(error) => {
362 self.stage.abort();
363 (Some(error), false, None)
364 }
365 Ok(output) => {
366 let (resume, terminal) =
367 self.durable.complete_with_terminal_response(job, output)?;
368 (None, resume, Some(terminal))
369 }
370 };
371 if self
372 .usage
373 .finish_inference(&mut self.durable, &key, terminal_id)
374 .await
375 .is_err()
376 {
377 return Ok(true);
378 }
379 if let Some(error) = error {
380 self.stage.fail(error.clone());
381 self.durable.fail(job, error, true);
382 return Ok(false);
383 }
384 self.shim = Some(shim);
385 if !resume && !self.stage.search_active() {
386 self.stage.clear_authorization(&mut self.durable);
387 }
388 Ok(false)
389 }
390
391 fn restart(
392 &mut self,
393 context: AccessContext,
394 profile_id: ProfileId,
395 policy: AccessPolicy,
396 ) -> Result<(), ActorError> {
397 let generation = self
398 .generation
399 .checked_add(1)
400 .ok_or(ActorError::NotRestartable)?;
401 self.stage
402 .restart(&mut self.durable, context, profile_id, policy)?;
403 self.generation = generation;
404 self.active_key = format!("{}#restart-{generation}", self.base_key);
405 self.shim = Some(new_shim(
406 self.adapter.clone(),
407 self.active_key.clone(),
408 self.sender.clone(),
409 ));
410 Ok(())
411 }
412}
413
414fn map_transition(error: TransitionError) -> ActorError {
415 match error {
416 TransitionError::Unauthorized => ActorError::Unauthorized,
417 TransitionError::NotStalled => ActorError::NotStalled,
418 TransitionError::NotRestartable => ActorError::NotRestartable,
419 TransitionError::Internal(_) => ActorError::Closed,
420 }
421}
422
423fn reply_transition(reply: Reply<()>, result: Result<(), TransitionError>) -> bool {
424 let stop = matches!(result, Err(TransitionError::Internal(_)));
425 answer(reply, result.map_err(map_transition), stop)
426}
427
428fn answer<T>(reply: Reply<T>, result: Result<T, ActorError>, stop: bool) -> bool {
429 let _ = reply.send(result);
430 stop
431}