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_persistence::Session;
5use kcode_k1_chat_state::USER_MESSAGE_TYPE;
6use kcode_k1_chat_thread_actor_channel::{
7    AcceptedBox, ActorError, ActorShim, BoxValue, Handle, Message, ProviderInput,
8    ProviderInputKind, Reply, Snapshot, StageReply, channel, new_shim,
9};
10use kcode_k1_chat_thread_durable_state::{
11    DurableThread, PreparedSteer, Status, ToolCallId, TransitionError,
12};
13use kcode_k1_chat_thread_web_search::{WebSearchAction, WebSearchToolEvent};
14use kcode_k1_codex_adapter::{Adapter, ShimOutput};
15use std::sync::Arc;
16use tokio::{sync::mpsc, task::JoinHandle};
17
18pub fn open(
19    adapter: Adapter,
20    key: impl Into<String>,
21    session: Session,
22    kmap: Arc<K1AccessKmap>,
23    web_search: WebSearchAction,
24) -> Result<Handle, String> {
25    let durable = DurableThread::recover(session, kmap)?;
26    let (handle, sender, receiver) = channel();
27    let base_key = key.into();
28    let active_key = base_key.clone();
29    let actor = Actor {
30        durable,
31        shim: Some(new_shim(
32            adapter.clone(),
33            active_key.clone(),
34            sender.clone(),
35        )),
36        adapter,
37        base_key,
38        active_key,
39        generation: 0,
40        sender,
41        receiver,
42        waiters: Vec::new(),
43        inference: None,
44        steering: None,
45        searches: Vec::new(),
46        search_epoch: 0,
47        stage_reply: None,
48        job: None,
49        model_input: None,
50        web_search,
51    };
52    tokio::spawn(actor.run());
53    Ok(handle)
54}
55
56struct Actor {
57    durable: DurableThread,
58    shim: Option<ActorShim>,
59    adapter: Adapter,
60    base_key: String,
61    active_key: String,
62    generation: u64,
63    sender: mpsc::UnboundedSender<Message>,
64    receiver: mpsc::UnboundedReceiver<Message>,
65    waiters: Vec<Reply<Snapshot>>,
66    inference: Option<JoinHandle<()>>,
67    steering: Option<JoinHandle<()>>,
68    searches: Vec<(ToolCallId, JoinHandle<()>)>,
69    search_epoch: u64,
70    stage_reply: Option<StageReply>,
71    job: Option<u64>,
72    model_input: Option<ProviderInput>,
73    web_search: WebSearchAction,
74}
75
76impl Actor {
77    async fn run(mut self) {
78        'actor: loop {
79            if self.drive() {
80                break;
81            }
82            self.wake();
83            let Some(mut message) = self.receiver.recv().await else {
84                break;
85            };
86            loop {
87                if self.handle(message) {
88                    break 'actor;
89                }
90                match self.receiver.try_recv() {
91                    Ok(next) => message = next,
92                    Err(_) => break,
93                }
94            }
95        }
96        self.abort_searches();
97        if let Some(reply) = self.stage_reply.take() {
98            let _ = reply.send(Err("K1 actor is closed".to_owned()));
99        }
100        for task in [self.steering.take(), self.inference.take()]
101            .into_iter()
102            .flatten()
103        {
104            task.abort();
105        }
106        for waiter in self.waiters {
107            let _ = waiter.send(Err(ActorError::Closed));
108        }
109    }
110
111    fn handle(&mut self, message: Message) -> bool {
112        match message {
113            Message::Accept(value, reply) => self.accept(value, reply),
114            Message::AcceptUser(context, profile_id, policy, contents, reply) => {
115                self.accept_user(context, profile_id, policy, contents, reply)
116            }
117            Message::Return(id, result, reply) => self.accept_return(id, result, reply),
118            Message::Stage(text, boxes, reply) => self.stage(text, boxes, reply),
119            Message::WebSearch(epoch, id, event) => self.web_search_event(epoch, id, event),
120            Message::WebSearchEnded(epoch, id) => self.web_search_ended(epoch, id),
121            Message::Steered(job, prepared, result) => self.steered(job, prepared, result),
122            Message::Inferred(job, shim, result) => self.finish(job, shim, result),
123            Message::Snapshot(reply) => {
124                let _ = reply.send(Ok(self.snapshot()));
125                false
126            }
127            Message::Wait(reply) => {
128                if self.running() {
129                    self.waiters.push(reply);
130                } else {
131                    let _ = reply.send(Ok(self.snapshot()));
132                }
133                false
134            }
135            Message::Restart(context, profile_id, policy, reply) => {
136                let result = self.restart(context, profile_id, policy);
137                let _ = reply.send(result);
138                false
139            }
140            Message::Abandon => true,
141        }
142    }
143
144    fn accept(&mut self, value: AcceptedBox, reply: Reply<()>) -> bool {
145        let (box_type, contents, hidden_type, hidden_contents) = value;
146        if box_type == USER_MESSAGE_TYPE {
147            let _ = reply.send(Err(ActorError::Unauthorized));
148            return false;
149        }
150        let result = self
151            .durable
152            .accept_box(box_type, contents, hidden_type, hidden_contents);
153        let stop = result.is_err();
154        let _ = reply.send(result.map_err(|_| ActorError::Closed));
155        stop
156    }
157
158    fn accept_user(
159        &mut self,
160        context: kcode_k1_chat_thread_durable_state::AccessContext,
161        profile_id: kcode_k1_chat_thread_durable_state::ProfileId,
162        policy: kcode_k1_chat_thread_durable_state::AccessPolicy,
163        contents: String,
164        reply: Reply<()>,
165    ) -> bool {
166        let result = self
167            .durable
168            .accept_user(context, profile_id, policy, contents);
169        let stop = matches!(result, Err(TransitionError::Internal(_)));
170        let _ = reply.send(result.map_err(map_transition));
171        stop
172    }
173
174    fn accept_return(
175        &mut self,
176        id: ToolCallId,
177        result: Result<String, String>,
178        reply: Reply<()>,
179    ) -> bool {
180        let result = self.durable.accept_return(id, result);
181        let stop = result.is_err();
182        let _ = reply.send(result.map_err(|_| ActorError::Closed));
183        stop
184    }
185
186    fn stage(&mut self, text: String, boxes: Vec<BoxValue>, reply: StageReply) -> bool {
187        if self.stage_reply.is_some() || self.steering.is_some() {
188            return reject_stage(reply, "overlapping Codex stages or steer".into());
189        }
190        let Some(job) = self.job else {
191            return reject_stage(reply, "no active K1 inference".into());
192        };
193        let calls = match self.durable.prepare_stage(job, text, boxes) {
194            Ok(calls) => calls,
195            Err(error) => return reject_stage(reply, error),
196        };
197        self.stage_reply = Some(reply);
198        for call in calls {
199            if call.name == self.web_search.tool_name() {
200                if self.search_live(&call.tool_call_id) {
201                    return self.fatal("duplicate WebSearch ToolCallId".into());
202                }
203                self.launch_search(call.tool_call_id, call.arguments);
204            } else {
205                let result = self.durable.launch_action(&call.name, &call.arguments);
206                if let Err(error) = self.durable.accept_tool_return(call.tool_call_id, result) {
207                    return self.fatal(error);
208                }
209            }
210        }
211        match self.start_steer() {
212            Ok(true) => self.finish_stage(),
213            Ok(false) => false,
214            Err(error) => self.fatal(error),
215        }
216    }
217
218    fn launch_search(&mut self, id: ToolCallId, arguments: String) {
219        let mut search = self.web_search.launch(&arguments);
220        let sender = self.sender.clone();
221        let epoch = self.search_epoch;
222        let event_id = id;
223        let task = tokio::spawn(async move {
224            while let Some(event) = search.recv().await {
225                let terminal = matches!(&event, WebSearchToolEvent::Result { .. });
226                if sender
227                    .send(Message::WebSearch(epoch, event_id, event))
228                    .is_err()
229                {
230                    return;
231                }
232                if terminal {
233                    return;
234                }
235            }
236            let _ = sender.send(Message::WebSearchEnded(epoch, event_id));
237        });
238        self.searches.push((id, task));
239    }
240
241    fn web_search_event(&mut self, epoch: u64, id: ToolCallId, event: WebSearchToolEvent) -> bool {
242        if epoch != self.search_epoch {
243            return false;
244        }
245        match event {
246            WebSearchToolEvent::Message { contents } => {
247                if !self.search_live(&id) {
248                    return self.fatal("WebSearch Message has no live task".into());
249                }
250                if let Err(error) = self.durable.accept_tool_message(id, contents) {
251                    return self.fatal(error);
252                }
253                false
254            }
255            WebSearchToolEvent::Result {
256                result,
257                metadata_type,
258                metadata_contents,
259            } => {
260                let Some(task) = self.remove_search(&id) else {
261                    return self.fatal("WebSearch Result has no live task".into());
262                };
263                task.abort();
264                if let Err(error) =
265                    self.durable
266                        .accept_tool_return_v2(id, result, metadata_type, metadata_contents)
267                {
268                    return self.fatal(error);
269                }
270                if self.job.is_none() {
271                    return false;
272                }
273                match self.start_steer() {
274                    Ok(_) => false,
275                    Err(error) => self.fatal(error),
276                }
277            }
278        }
279    }
280
281    fn web_search_ended(&mut self, epoch: u64, id: ToolCallId) -> bool {
282        if epoch != self.search_epoch {
283            return false;
284        }
285        let Some(task) = self.remove_search(&id) else {
286            return self.fatal("ended WebSearch has no live task".into());
287        };
288        task.abort();
289        self.fatal("WebSearch ended without a terminal Result".into())
290    }
291
292    fn start_steer(&mut self) -> Result<bool, String> {
293        let job = self
294            .job
295            .ok_or_else(|| "no active K1 inference".to_owned())?;
296        let prepared = self.durable.prepare_steer(job)?;
297        if self.steering.is_some() {
298            return Ok(false);
299        }
300        let Some(prepared) = prepared else {
301            return Ok(true);
302        };
303        let input = self.durable.prepared_input(&prepared)?;
304        self.model_input = Some(ProviderInput {
305            kind: ProviderInputKind::Steer,
306            text: input.clone(),
307        });
308        let adapter = self.adapter.clone();
309        let key = self.active_key.clone();
310        let sender = self.sender.clone();
311        self.steering = Some(tokio::spawn(async move {
312            let result = adapter
313                .steer(key, input)
314                .await
315                .map_err(|error| error.to_string());
316            let _ = sender.send(Message::Steered(job, prepared, result));
317        }));
318        Ok(false)
319    }
320
321    fn steered(&mut self, job: u64, prepared: PreparedSteer, result: Result<(), String>) -> bool {
322        if self.steering.take().is_none() || self.job != Some(job) {
323            return self.fatal("stale Codex steer completion".into());
324        }
325        if let Err(error) = result {
326            return self.fatal(error);
327        }
328        if let Err(error) = self.durable.commit_steer(prepared) {
329            return self.fatal(error);
330        }
331        if self.stage_reply.is_some() && self.finish_stage() {
332            return true;
333        }
334        match self.start_steer() {
335            Ok(_) => false,
336            Err(error) => self.fatal(error),
337        }
338    }
339
340    fn finish_stage(&mut self) -> bool {
341        match self.stage_reply.take() {
342            Some(reply) => reply.send(Ok(())).is_err(),
343            None => false,
344        }
345    }
346
347    fn fatal(&mut self, error: String) -> bool {
348        if let Some(reply) = self.stage_reply.take() {
349            let _ = reply.send(Err(error));
350        }
351        true
352    }
353
354    fn search_live(&self, id: &ToolCallId) -> bool {
355        self.searches.iter().any(|(live, _)| live == id)
356    }
357
358    fn remove_search(&mut self, id: &ToolCallId) -> Option<JoinHandle<()>> {
359        let index = self.searches.iter().position(|(live, _)| live == id)?;
360        Some(self.searches.swap_remove(index).1)
361    }
362
363    fn abort_searches(&mut self) {
364        for (_, task) in self.searches.drain(..) {
365            task.abort();
366        }
367    }
368
369    fn cancel_searches(&mut self) -> Result<(), String> {
370        self.abort_searches();
371        self.search_epoch = self
372            .search_epoch
373            .checked_add(1)
374            .ok_or_else(|| "WebSearch cancellation epoch was exhausted".to_owned())?;
375        Ok(())
376    }
377
378    fn drive(&mut self) -> bool {
379        if self.inference.is_some() || self.shim.is_none() {
380            return false;
381        }
382        let (job, input) = match self.durable.begin_input() {
383            Ok(Some(value)) => value,
384            Ok(None) => return false,
385            Err(error) => return self.fatal(error),
386        };
387        self.model_input = Some(ProviderInput {
388            kind: ProviderInputKind::Turn,
389            text: input.clone(),
390        });
391        let mut shim = self.shim.take().expect("shim was checked");
392        self.job = Some(job);
393        let sender = self.sender.clone();
394        self.inference = Some(tokio::spawn(async move {
395            let result = shim.infer(input).await.map_err(|error| error.to_string());
396            let _ = sender.send(Message::Inferred(job, shim, result));
397        }));
398        false
399    }
400
401    fn finish(
402        &mut self,
403        job: u64,
404        shim: ActorShim,
405        result: Result<ShimOutput<BoxValue>, String>,
406    ) -> bool {
407        if self.inference.take().is_none() || self.steering.is_some() || self.job != Some(job) {
408            return self.fatal("stale Codex inference completion".into());
409        }
410        self.job = None;
411        match result {
412            Err(error) => {
413                let fence = self.cancel_searches();
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                match fence {
419                    Ok(()) => false,
420                    Err(error) => self.fatal(error),
421                }
422            }
423            Ok(output) => {
424                let resume = match self.durable.complete(job, output) {
425                    Ok(resume) => resume,
426                    Err(error) => return self.fatal(error),
427                };
428                self.shim = Some(shim);
429                if !resume && self.searches.is_empty() {
430                    self.durable.clear_authorization();
431                }
432                false
433            }
434        }
435    }
436
437    fn restart(
438        &mut self,
439        context: kcode_k1_chat_thread_durable_state::AccessContext,
440        profile_id: kcode_k1_chat_thread_durable_state::ProfileId,
441        policy: kcode_k1_chat_thread_durable_state::AccessPolicy,
442    ) -> Result<(), ActorError> {
443        let generation = self
444            .generation
445            .checked_add(1)
446            .ok_or(ActorError::NotRestartable)?;
447        self.durable
448            .restart(context, profile_id, policy)
449            .map_err(map_transition)?;
450        self.generation = generation;
451        self.active_key = format!("{}#restart-{generation}", self.base_key);
452        self.shim = Some(new_shim(
453            self.adapter.clone(),
454            self.active_key.clone(),
455            self.sender.clone(),
456        ));
457        Ok(())
458    }
459
460    fn running(&self) -> bool {
461        !self.searches.is_empty() || matches!(self.durable.status(), Status::Running)
462    }
463
464    fn snapshot(&self) -> Snapshot {
465        let status = if self.searches.is_empty() {
466            self.durable.status()
467        } else {
468            Status::Running
469        };
470        Snapshot {
471            boxes: self.durable.boxes().to_vec(),
472            status,
473            model_input: self.model_input.clone(),
474        }
475    }
476
477    fn wake(&mut self) {
478        if !self.running() {
479            let snapshot = self.snapshot();
480            for waiter in std::mem::take(&mut self.waiters) {
481                let _ = waiter.send(Ok(snapshot.clone()));
482            }
483        }
484    }
485}
486
487fn map_transition(error: TransitionError) -> ActorError {
488    match error {
489        TransitionError::Unauthorized => ActorError::Unauthorized,
490        TransitionError::NotStalled => ActorError::NotStalled,
491        TransitionError::NotRestartable => ActorError::NotRestartable,
492        TransitionError::Internal(_) => ActorError::Closed,
493    }
494}
495
496fn reject_stage(reply: StageReply, error: String) -> bool {
497    let _ = reply.send(Err(error));
498    true
499}