Skip to main content

kcode_k1_chat_thread_session_inference_settlement/
lib.rs

1#![forbid(unsafe_code)]
2#![doc = include_str!("../Documentation.md")]
3
4use kcode_k1_chat_model_usage_session::ModelUsageSession;
5use kcode_k1_chat_thread_actor_channel::{ActorShim, BoxValue};
6use kcode_k1_chat_thread_durable_state::{ChatDiagnostic, DurableThread};
7use kcode_k1_chat_thread_rust_code_session::{PersistCompletion, RustCodeSession};
8use kcode_k1_chat_thread_session_provider::SessionProvider;
9use kcode_k1_chat_thread_session_stage_runtime::StageRuntime;
10use kcode_k1_chat_thread_web_search_tasks::WebSearchTasks;
11use kcode_k1_codex_adapter::ShimOutput;
12
13const SAFE_RECOVERABLE_FAILURE: &str = "The response could not be completed. Please try again.";
14
15pub struct InferenceSettlement {
16    stop: bool,
17    stage_error: Option<String>,
18    recoverable_failure: bool,
19    usage_diagnostic: Option<ChatDiagnostic>,
20}
21
22impl InferenceSettlement {
23    pub fn should_stop(&self) -> bool {
24        self.stop
25    }
26
27    pub fn recoverable_failure(&self) -> bool {
28        self.recoverable_failure
29    }
30
31    pub fn usage_diagnostic(&self) -> Option<ChatDiagnostic> {
32        self.usage_diagnostic
33    }
34
35    pub fn into_stage_error(self) -> Option<String> {
36        self.stage_error
37    }
38}
39
40pub struct StageInferenceSettlement {
41    stop: bool,
42    restore_shim: bool,
43    stage_error: Option<String>,
44    recoverable_failure: bool,
45    usage_diagnostic: Option<ChatDiagnostic>,
46}
47
48impl StageInferenceSettlement {
49    pub fn should_stop(&self) -> bool {
50        self.stop
51    }
52
53    pub fn restore_shim(&self) -> bool {
54        self.restore_shim
55    }
56
57    pub fn recoverable_failure(&self) -> bool {
58        self.recoverable_failure
59    }
60
61    pub fn usage_diagnostic(&self) -> Option<ChatDiagnostic> {
62        self.usage_diagnostic
63    }
64
65    pub fn into_stage_error(self) -> Option<String> {
66        self.stage_error
67    }
68}
69
70#[allow(clippy::too_many_arguments)]
71pub async fn settle_inference(
72    durable: &mut DurableThread,
73    provider: &mut SessionProvider,
74    usage: &mut ModelUsageSession,
75    searches: &mut WebSearchTasks,
76    mut rust_code: Option<&mut RustCodeSession>,
77    job: u64,
78    shim: ActorShim,
79    result: Result<ShimOutput<BoxValue>, String>,
80) -> Result<InferenceSettlement, String> {
81    provider.accept_inferred(job)?;
82    let key = provider.active_key().to_owned();
83    match result {
84        Err(_) => {
85            searches.abort();
86            if let Some(rust) = rust_code.as_deref_mut() {
87                rust.abort();
88            }
89            let (_, terminal) = durable.complete_recoverable_failure_with_terminal_response(
90                job,
91                SAFE_RECOVERABLE_FAILURE.to_owned(),
92            )?;
93            durable.reset_provider_context()?;
94            let _ = durable.record_diagnostic(ChatDiagnostic::ProviderInference);
95            if let Some(rust) = rust_code.as_deref_mut() {
96                rust.clear_authorization_mirror();
97            }
98            provider.settle_failure();
99            let usage_diagnostic = finish_usage(usage, durable, &key, Some(terminal), true).await;
100            Ok(InferenceSettlement {
101                stop: false,
102                stage_error: Some(SAFE_RECOVERABLE_FAILURE.to_owned()),
103                recoverable_failure: true,
104                usage_diagnostic,
105            })
106        }
107        Ok(output) => {
108            let active_without_completion = rust_code
109                .as_ref()
110                .is_some_and(|rust| rust.is_busy() && !rust.has_completion());
111            let (resume, terminal) = durable.complete_with_terminal_response(job, output)?;
112            if rust_code.as_ref().is_some_and(|rust| rust.has_completion()) {
113                let outcome = rust_code
114                    .as_deref_mut()
115                    .expect("completion requires Rust support")
116                    .persist_completion(durable, job)?;
117                if !matches!(outcome, PersistCompletion::Committed(None)) {
118                    return Err(
119                        "Rust code completion did not commit at terminal boundary".to_owned()
120                    );
121                }
122            }
123            if active_without_completion {
124                return Err("Codex inference completed while Rust code work remains".to_owned());
125            }
126            if rust_code
127                .as_ref()
128                .is_some_and(|rust| rust.has_completion() || rust.is_busy())
129            {
130                return Err("Rust code work restarted at terminal boundary".to_owned());
131            }
132            let usage_diagnostic = finish_usage(usage, durable, &key, Some(terminal), false).await;
133            provider.settle_success(shim);
134            if !resume && !searches.is_busy() {
135                if let Some(rust) = rust_code {
136                    rust.clear_authorization(durable);
137                } else {
138                    durable.clear_authorization();
139                }
140            }
141            Ok(InferenceSettlement {
142                stop: false,
143                stage_error: None,
144                recoverable_failure: false,
145                usage_diagnostic,
146            })
147        }
148    }
149}
150
151pub async fn settle_stage_inference(
152    durable: &mut DurableThread,
153    usage: &mut ModelUsageSession,
154    stage: &mut StageRuntime,
155    active_key: &str,
156    job: u64,
157    result: Result<ShimOutput<BoxValue>, String>,
158) -> Result<StageInferenceSettlement, String> {
159    match result {
160        Err(_) => {
161            settle_stage_failure(
162                durable,
163                usage,
164                stage,
165                active_key,
166                job,
167                ChatDiagnostic::ProviderInference,
168            )
169            .await
170        }
171        Ok(output) => {
172            let (resume, terminal) = durable.complete_with_terminal_response(job, output)?;
173            stage.settle_terminal_code(durable, job)?;
174            let usage_diagnostic =
175                finish_usage(usage, durable, active_key, Some(terminal), false).await;
176            if !resume && !stage.search_active() {
177                stage.clear_authorization(durable);
178            }
179            Ok(StageInferenceSettlement {
180                stop: false,
181                restore_shim: true,
182                stage_error: None,
183                recoverable_failure: false,
184                usage_diagnostic,
185            })
186        }
187    }
188}
189
190pub async fn settle_stage_failure(
191    durable: &mut DurableThread,
192    usage: &mut ModelUsageSession,
193    stage: &mut StageRuntime,
194    active_key: &str,
195    job: u64,
196    diagnostic: ChatDiagnostic,
197) -> Result<StageInferenceSettlement, String> {
198    stage.abort();
199    let (_, terminal) = durable.complete_recoverable_failure_with_terminal_response(
200        job,
201        SAFE_RECOVERABLE_FAILURE.to_owned(),
202    )?;
203    durable.reset_provider_context()?;
204    let _ = durable.record_diagnostic(diagnostic);
205    stage.clear_authorization(durable);
206    let usage_diagnostic = finish_usage(usage, durable, active_key, Some(terminal), true).await;
207    Ok(StageInferenceSettlement {
208        stop: false,
209        restore_shim: false,
210        stage_error: Some(SAFE_RECOVERABLE_FAILURE.to_owned()),
211        recoverable_failure: true,
212        usage_diagnostic,
213    })
214}
215
216async fn finish_usage(
217    usage: &mut ModelUsageSession,
218    durable: &mut DurableThread,
219    active_key: &str,
220    terminal: Option<u64>,
221    restart: bool,
222) -> Option<ChatDiagnostic> {
223    if !usage.is_subscribed() {
224        return None;
225    }
226    if usage
227        .finish_inference(durable, active_key, terminal)
228        .await
229        .is_err()
230    {
231        return Some(ChatDiagnostic::ModelUsageFinish);
232    }
233    if restart && usage.restart(durable).await.is_err() {
234        return Some(ChatDiagnostic::ModelUsageRestart);
235    }
236    None
237}