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}