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}