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