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_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}