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