Skip to main content

kcode_k1_chat_thread_session_actor/
lib.rs

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