1#![forbid(unsafe_code)]
2
3use kcode_k1_access_kmap::K1AccessKmap;
4use kcode_k1_chat_codex_state::PreparedCallDisposition;
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 StageReply, channel_with_events, new_shim,
10};
11use kcode_k1_chat_thread_durable_state::{
12 DurableThread, PreparedMailboxFlush, ToolCallId, TransitionError,
13};
14use kcode_k1_chat_thread_session_view::SessionView;
15use kcode_k1_chat_thread_web_search::{WebSearchAction, WebSearchToolEvent};
16use kcode_k1_chat_thread_web_search_tasks::WebSearchTasks;
17use kcode_k1_codex_adapter::{Adapter, ShimOutput};
18use std::sync::Arc;
19use tokio::{sync::mpsc, task::JoinHandle};
20
21pub fn open(
22 adapter: Adapter,
23 key: impl Into<String>,
24 session: Session,
25 kmap: Arc<K1AccessKmap>,
26 web_search: WebSearchAction,
27) -> Result<Handle, String> {
28 let durable = DurableThread::recover(session, kmap)?;
29 let (handle, sender, receiver, event_receiver) = channel_with_events();
30 let base_key = key.into();
31 let active_key = base_key.clone();
32 let actor = Actor {
33 durable,
34 shim: Some(new_shim(
35 adapter.clone(),
36 active_key.clone(),
37 sender.clone(),
38 )),
39 adapter,
40 base_key,
41 active_key,
42 generation: 0,
43 web_search: WebSearchTasks::new(web_search, sender.clone()),
44 view: SessionView::new(event_receiver),
45 sender,
46 receiver,
47 usage: ModelUsageSession::default(),
48 inference: None,
49 mailbox_flush: None,
50 stage_reply: None,
51 job: None,
52 };
53 tokio::spawn(actor.run());
54 Ok(handle)
55}
56
57struct Actor {
58 durable: DurableThread,
59 shim: Option<ActorShim>,
60 adapter: Adapter,
61 base_key: String,
62 active_key: String,
63 generation: u64,
64 web_search: WebSearchTasks,
65 view: SessionView,
66 sender: mpsc::UnboundedSender<Message>,
67 receiver: mpsc::UnboundedReceiver<Message>,
68 usage: ModelUsageSession,
69 inference: Option<JoinHandle<()>>,
70 mailbox_flush: Option<JoinHandle<()>>,
71 stage_reply: Option<StageReply>,
72 job: Option<u64>,
73}
74
75impl Actor {
76 async fn run(mut self) {
77 'actor: loop {
78 if self.drive().await {
79 break;
80 }
81 self.view.wake(&self.durable, !self.web_search.is_empty());
82 tokio::select! {
83 message = self.receiver.recv() => {
84 let Some(mut message) = message else { break; };
85 loop {
86 if self.handle(message).await {
87 break 'actor;
88 }
89 match self.receiver.try_recv() {
90 Ok(next) => message = next,
91 Err(_) => break,
92 }
93 }
94 }
95 reply = self.view.receive_event_query() => {
96 if let Some(reply) = reply {
97 self.view.answer_event_query(
98 &self.durable,
99 !self.web_search.is_empty(),
100 reply,
101 );
102 }
103 }
104 usage = self.usage.receive(&mut self.durable) => {
105 if usage.is_err() {
106 break;
107 }
108 }
109 }
110 }
111 let _ = self
112 .usage
113 .shutdown(&mut self.durable, self.inference.is_some())
114 .await;
115 self.web_search.abort();
116 if let Some(reply) = self.stage_reply.take() {
117 let _ = reply.send(Err("K1 actor is closed".to_owned()));
118 }
119 for task in [self.mailbox_flush.take(), self.inference.take()]
120 .into_iter()
121 .flatten()
122 {
123 task.abort();
124 }
125 self.view.close();
126 }
127
128 async fn handle(&mut self, message: Message) -> bool {
129 match message {
130 Message::Accept((box_type, contents, hidden_type, hidden_contents), reply) => {
131 let result = self.durable.accept_external_box(
132 box_type,
133 contents,
134 hidden_type,
135 hidden_contents,
136 );
137 let stop = matches!(result, Err(TransitionError::Internal(_)));
138 let _ = reply.send(result.map_err(map_transition));
139 stop
140 }
141 Message::AcceptUser(context, profile_id, policy, contents, reply) => {
142 self.accept_user(context, profile_id, policy, contents, reply)
143 }
144 Message::Return(id, result, reply) => self.accept_return(id, result, reply),
145 Message::Stage(text, boxes, reply) => self.stage(text, boxes, reply),
146 Message::WebSearch(epoch, id, event) => self.web_search_event(epoch, id, event),
147 Message::WebSearchEnded(epoch, id) => self.web_search_ended(epoch, id),
148 Message::MailboxFlushCompleted(job, prepared, result) => {
149 self.mailbox_flush_completed(job, prepared, result)
150 }
151 Message::Inferred(job, shim, result) => self.finish(job, shim, result).await,
152 Message::Snapshot(reply) => {
153 let snapshot = self
154 .view
155 .snapshot(&self.durable, !self.web_search.is_empty());
156 let _ = reply.send(Ok(snapshot));
157 false
158 }
159 Message::Wait(reply) => {
160 self.view
161 .wait(&self.durable, !self.web_search.is_empty(), reply);
162 false
163 }
164 Message::Restart(context, profile_id, policy, reply) => {
165 if self.usage.restart(&mut self.durable).await.is_err() {
166 let _ = reply.send(Err(ActorError::Closed));
167 true
168 } else {
169 let result = self.restart(context, profile_id, policy);
170 let _ = reply.send(result);
171 false
172 }
173 }
174 Message::Abandon => true,
175 }
176 }
177
178 fn accept_user(
179 &mut self,
180 context: kcode_k1_chat_thread_durable_state::AccessContext,
181 profile_id: kcode_k1_chat_thread_durable_state::ProfileId,
182 policy: kcode_k1_chat_thread_durable_state::AccessPolicy,
183 contents: String,
184 reply: Reply<()>,
185 ) -> bool {
186 let result = self
187 .durable
188 .accept_user(context, profile_id, policy, contents);
189 let stop = matches!(result, Err(TransitionError::Internal(_)));
190 let _ = reply.send(result.map_err(map_transition));
191 stop
192 }
193
194 fn accept_return(
195 &mut self,
196 id: ToolCallId,
197 result: Result<String, String>,
198 reply: Reply<()>,
199 ) -> bool {
200 let result = self.durable.accept_return(id, result);
201 let stop = result.is_err();
202 let _ = reply.send(result.map_err(|_| ActorError::Closed));
203 stop
204 }
205
206 fn stage(&mut self, text: String, boxes: Vec<BoxValue>, reply: StageReply) -> bool {
207 if self.stage_reply.is_some() || self.mailbox_flush.is_some() {
208 return reject_stage(reply, "overlapping Codex stages or mailbox flush".into());
209 }
210 let Some(job) = self.job else {
211 return reject_stage(reply, "no active K1 inference".into());
212 };
213 let calls = match self.durable.prepare_stage(job, text, boxes) {
214 Ok(calls) => calls,
215 Err(error) => return reject_stage(reply, error),
216 };
217 self.stage_reply = Some(reply);
218 for call in calls {
219 match call.disposition().clone() {
220 PreparedCallDisposition::ImmediateError(error) => {
221 if let Err(error) = self.durable.accept_tool_return_v2(
222 call.tool_call_id,
223 Err(error.message.clone()),
224 error.metadata_type.clone(),
225 error.metadata_contents.clone(),
226 ) {
227 return self.fatal(error);
228 }
229 }
230 PreparedCallDisposition::External => {
231 if call.name == self.web_search.tool_name() {
232 if let Err(error) =
233 self.web_search.launch(call.tool_call_id, call.arguments)
234 {
235 return self.fatal(error);
236 }
237 } else {
238 let result = self.durable.launch_action(&call.name, &call.arguments);
239 if let Err(error) =
240 self.durable.accept_tool_return(call.tool_call_id, result)
241 {
242 return self.fatal(error);
243 }
244 }
245 }
246 }
247 }
248 match self.start_mailbox_flush() {
249 Ok(true) => self.finish_stage(),
250 Ok(false) => false,
251 Err(error) => self.fatal(error),
252 }
253 }
254
255 fn web_search_event(&mut self, epoch: u64, id: ToolCallId, event: WebSearchToolEvent) -> bool {
256 let event = match self.web_search.accept_event(epoch, id, event) {
257 Ok(Some(event)) => event,
258 Ok(None) => return false,
259 Err(error) => return self.fatal(error),
260 };
261 match event {
262 WebSearchToolEvent::Message { contents } => {
263 if let Err(error) = self.durable.accept_tool_message(id, contents) {
264 return self.fatal(error);
265 }
266 false
267 }
268 WebSearchToolEvent::Result {
269 result,
270 metadata_type,
271 metadata_contents,
272 } => {
273 if let Err(error) =
274 self.durable
275 .accept_tool_return_v2(id, result, metadata_type, metadata_contents)
276 {
277 return self.fatal(error);
278 }
279 if self.job.is_none() {
280 return false;
281 }
282 match self.start_mailbox_flush() {
283 Ok(_) => false,
284 Err(error) => self.fatal(error),
285 }
286 }
287 }
288 }
289
290 fn web_search_ended(&mut self, epoch: u64, id: ToolCallId) -> bool {
291 match self.web_search.accept_ended(epoch, id) {
292 Ok(()) => false,
293 Err(error) => self.fatal(error),
294 }
295 }
296
297 fn start_mailbox_flush(&mut self) -> Result<bool, String> {
298 let job = self
299 .job
300 .ok_or_else(|| "no active K1 inference".to_owned())?;
301 let prepared = self.durable.prepare_mailbox_flush(job)?;
302 if self.mailbox_flush.is_some() {
303 return Ok(false);
304 }
305 let Some(prepared) = prepared else {
306 return Ok(true);
307 };
308 let input = self.durable.prepared_input(&prepared)?;
309 self.view.set_model_input(ProviderInput {
310 kind: ProviderInputKind::MailboxFlush,
311 text: input.clone(),
312 });
313 let adapter = self.adapter.clone();
314 let key = self.active_key.clone();
315 let sender = self.sender.clone();
316 self.mailbox_flush = Some(tokio::spawn(async move {
317 let result = adapter
318 .steer(key, input)
319 .await
320 .map_err(|error| error.to_string());
321 let _ = sender.send(Message::MailboxFlushCompleted(job, prepared, result));
322 }));
323 Ok(false)
324 }
325
326 fn mailbox_flush_completed(
327 &mut self,
328 job: u64,
329 prepared: PreparedMailboxFlush,
330 result: Result<(), String>,
331 ) -> bool {
332 if self.mailbox_flush.take().is_none() || self.job != Some(job) {
333 return self.fatal("stale Codex mailbox-flush completion".into());
334 }
335 if let Err(error) = result {
336 return self.fatal(error);
337 }
338 if let Err(error) = self.durable.commit_mailbox_flush(prepared) {
339 return self.fatal(error);
340 }
341 if self.stage_reply.is_some() && self.finish_stage() {
342 return true;
343 }
344 match self.start_mailbox_flush() {
345 Ok(_) => false,
346 Err(error) => self.fatal(error),
347 }
348 }
349
350 fn finish_stage(&mut self) -> bool {
351 match self.stage_reply.take() {
352 Some(reply) => reply.send(Ok(())).is_err(),
353 None => false,
354 }
355 }
356
357 fn fatal(&mut self, error: String) -> bool {
358 if let Some(reply) = self.stage_reply.take() {
359 let _ = reply.send(Err(error));
360 }
361 true
362 }
363
364 async fn drive(&mut self) -> bool {
365 if self.inference.is_some() || self.shim.is_none() {
366 return false;
367 }
368 let (job, input) = match self.durable.begin_input() {
369 Ok(Some(value)) => value,
370 Ok(None) => return false,
371 Err(error) => return self.fatal(error),
372 };
373 if !self.usage.is_subscribed()
374 && self
375 .usage
376 .subscribe(&self.adapter, self.active_key.clone())
377 .await
378 .is_err()
379 {
380 return self.fatal("model-usage subscription failed".into());
381 }
382 self.view.set_model_input(ProviderInput {
383 kind: ProviderInputKind::Turn,
384 text: input.clone(),
385 });
386 let mut shim = self.shim.take().expect("shim was checked");
387 self.job = Some(job);
388 let sender = self.sender.clone();
389 self.inference = Some(tokio::spawn(async move {
390 let result = shim.infer(input).await.map_err(|error| error.to_string());
391 let _ = sender.send(Message::Inferred(job, shim, result));
392 }));
393 false
394 }
395
396 async fn finish(
397 &mut self,
398 job: u64,
399 shim: ActorShim,
400 result: Result<ShimOutput<BoxValue>, String>,
401 ) -> bool {
402 if self.inference.take().is_none() || self.mailbox_flush.is_some() || self.job != Some(job)
403 {
404 return self.fatal("stale Codex inference completion".into());
405 }
406 self.job = None;
407 let key = self.active_key.clone();
408 match result {
409 Err(error) => {
410 if self
411 .usage
412 .finish_inference(&mut self.durable, &key, None)
413 .await
414 .is_err()
415 {
416 return true;
417 }
418 let fence = self.web_search.cancel();
419 if let Some(reply) = self.stage_reply.take() {
420 let _ = reply.send(Err(error.clone()));
421 }
422 self.durable.fail(job, error, true);
423 match fence {
424 Ok(()) => false,
425 Err(error) => self.fatal(error),
426 }
427 }
428 Ok(output) => {
429 let (resume, terminal_id) =
430 match self.durable.complete_with_terminal_response(job, output) {
431 Ok(value) => value,
432 Err(error) => return self.fatal(error),
433 };
434 if self
435 .usage
436 .finish_inference(&mut self.durable, &key, Some(terminal_id))
437 .await
438 .is_err()
439 {
440 return true;
441 }
442 self.shim = Some(shim);
443 if !resume && self.web_search.is_empty() {
444 self.durable.clear_authorization();
445 }
446 false
447 }
448 }
449 }
450
451 fn restart(
452 &mut self,
453 context: kcode_k1_chat_thread_durable_state::AccessContext,
454 profile_id: kcode_k1_chat_thread_durable_state::ProfileId,
455 policy: kcode_k1_chat_thread_durable_state::AccessPolicy,
456 ) -> Result<(), ActorError> {
457 let generation = self
458 .generation
459 .checked_add(1)
460 .ok_or(ActorError::NotRestartable)?;
461 self.durable
462 .restart(context, profile_id, policy)
463 .map_err(map_transition)?;
464 self.generation = generation;
465 self.active_key = format!("{}#restart-{generation}", self.base_key);
466 self.shim = Some(new_shim(
467 self.adapter.clone(),
468 self.active_key.clone(),
469 self.sender.clone(),
470 ));
471 Ok(())
472 }
473}
474
475fn map_transition(error: TransitionError) -> ActorError {
476 match error {
477 TransitionError::Unauthorized => ActorError::Unauthorized,
478 TransitionError::NotStalled => ActorError::NotStalled,
479 TransitionError::NotRestartable => ActorError::NotRestartable,
480 TransitionError::Internal(_) => ActorError::Closed,
481 }
482}
483
484fn reject_stage(reply: StageReply, error: String) -> bool {
485 let _ = reply.send(Err(error));
486 true
487}