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