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