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