Skip to main content

kcode_k1_chat_thread_session_actor/
lib.rs

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}