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