#![forbid(unsafe_code)]
#![doc = include_str!("../Documentation.md")]
use kcode_k1_chat_model_usage_session::ModelUsageSession;
use kcode_k1_chat_thread_actor_channel::{ActorShim, BoxValue};
use kcode_k1_chat_thread_durable_state::{ChatDiagnostic, DurableThread};
use kcode_k1_chat_thread_rust_code_session::{PersistCompletion, RustCodeSession};
use kcode_k1_chat_thread_session_provider::SessionProvider;
use kcode_k1_chat_thread_session_stage_runtime::StageRuntime;
use kcode_k1_chat_thread_web_search_tasks::WebSearchTasks;
use kcode_k1_codex_adapter::ShimOutput;
const SAFE_RECOVERABLE_FAILURE: &str = "The response could not be completed. Please try again.";
pub struct InferenceSettlement {
stop: bool,
stage_error: Option<String>,
recoverable_failure: bool,
usage_diagnostic: Option<ChatDiagnostic>,
}
impl InferenceSettlement {
pub fn should_stop(&self) -> bool {
self.stop
}
pub fn recoverable_failure(&self) -> bool {
self.recoverable_failure
}
pub fn usage_diagnostic(&self) -> Option<ChatDiagnostic> {
self.usage_diagnostic
}
pub fn into_stage_error(self) -> Option<String> {
self.stage_error
}
}
pub struct StageInferenceSettlement {
stop: bool,
restore_shim: bool,
stage_error: Option<String>,
recoverable_failure: bool,
usage_diagnostic: Option<ChatDiagnostic>,
}
impl StageInferenceSettlement {
pub fn should_stop(&self) -> bool {
self.stop
}
pub fn restore_shim(&self) -> bool {
self.restore_shim
}
pub fn recoverable_failure(&self) -> bool {
self.recoverable_failure
}
pub fn usage_diagnostic(&self) -> Option<ChatDiagnostic> {
self.usage_diagnostic
}
pub fn into_stage_error(self) -> Option<String> {
self.stage_error
}
}
#[allow(clippy::too_many_arguments)]
pub async fn settle_inference(
durable: &mut DurableThread,
provider: &mut SessionProvider,
usage: &mut ModelUsageSession,
searches: &mut WebSearchTasks,
mut rust_code: Option<&mut RustCodeSession>,
job: u64,
shim: ActorShim,
result: Result<ShimOutput<BoxValue>, String>,
) -> Result<InferenceSettlement, String> {
provider.accept_inferred(job)?;
let key = provider.active_key().to_owned();
match result {
Err(_) => {
searches.abort();
if let Some(rust) = rust_code.as_deref_mut() {
rust.abort();
}
let (_, terminal) = durable.complete_recoverable_failure_with_terminal_response(
job,
SAFE_RECOVERABLE_FAILURE.to_owned(),
)?;
durable.reset_provider_context()?;
let _ = durable.record_diagnostic(ChatDiagnostic::ProviderInference);
if let Some(rust) = rust_code.as_deref_mut() {
rust.clear_authorization_mirror();
}
provider.settle_failure();
let usage_diagnostic = finish_usage(usage, durable, &key, Some(terminal), true).await;
Ok(InferenceSettlement {
stop: false,
stage_error: Some(SAFE_RECOVERABLE_FAILURE.to_owned()),
recoverable_failure: true,
usage_diagnostic,
})
}
Ok(output) => {
let active_without_completion = rust_code
.as_ref()
.is_some_and(|rust| rust.is_busy() && !rust.has_completion());
let (resume, terminal) = durable.complete_with_terminal_response(job, output)?;
if rust_code.as_ref().is_some_and(|rust| rust.has_completion()) {
let outcome = rust_code
.as_deref_mut()
.expect("completion requires Rust support")
.persist_completion(durable, job)?;
if !matches!(outcome, PersistCompletion::Committed(None)) {
return Err(
"Rust code completion did not commit at terminal boundary".to_owned()
);
}
}
if active_without_completion {
return Err("Codex inference completed while Rust code work remains".to_owned());
}
if rust_code
.as_ref()
.is_some_and(|rust| rust.has_completion() || rust.is_busy())
{
return Err("Rust code work restarted at terminal boundary".to_owned());
}
let usage_diagnostic = finish_usage(usage, durable, &key, Some(terminal), false).await;
provider.settle_success(shim);
if !resume && !searches.is_busy() {
if let Some(rust) = rust_code {
rust.clear_authorization(durable);
} else {
durable.clear_authorization();
}
}
Ok(InferenceSettlement {
stop: false,
stage_error: None,
recoverable_failure: false,
usage_diagnostic,
})
}
}
}
pub async fn settle_stage_inference(
durable: &mut DurableThread,
usage: &mut ModelUsageSession,
stage: &mut StageRuntime,
active_key: &str,
job: u64,
result: Result<ShimOutput<BoxValue>, String>,
) -> Result<StageInferenceSettlement, String> {
match result {
Err(_) => {
settle_stage_failure(
durable,
usage,
stage,
active_key,
job,
ChatDiagnostic::ProviderInference,
)
.await
}
Ok(output) => {
let (resume, terminal) = durable.complete_with_terminal_response(job, output)?;
stage.settle_terminal_code(durable, job)?;
let usage_diagnostic =
finish_usage(usage, durable, active_key, Some(terminal), false).await;
if !resume && !stage.search_active() {
stage.clear_authorization(durable);
}
Ok(StageInferenceSettlement {
stop: false,
restore_shim: true,
stage_error: None,
recoverable_failure: false,
usage_diagnostic,
})
}
}
}
pub async fn settle_stage_failure(
durable: &mut DurableThread,
usage: &mut ModelUsageSession,
stage: &mut StageRuntime,
active_key: &str,
job: u64,
diagnostic: ChatDiagnostic,
) -> Result<StageInferenceSettlement, String> {
stage.abort();
let (_, terminal) = durable.complete_recoverable_failure_with_terminal_response(
job,
SAFE_RECOVERABLE_FAILURE.to_owned(),
)?;
durable.reset_provider_context()?;
let _ = durable.record_diagnostic(diagnostic);
stage.clear_authorization(durable);
let usage_diagnostic = finish_usage(usage, durable, active_key, Some(terminal), true).await;
Ok(StageInferenceSettlement {
stop: false,
restore_shim: false,
stage_error: Some(SAFE_RECOVERABLE_FAILURE.to_owned()),
recoverable_failure: true,
usage_diagnostic,
})
}
async fn finish_usage(
usage: &mut ModelUsageSession,
durable: &mut DurableThread,
active_key: &str,
terminal: Option<u64>,
restart: bool,
) -> Option<ChatDiagnostic> {
if !usage.is_subscribed() {
return None;
}
if usage
.finish_inference(durable, active_key, terminal)
.await
.is_err()
{
return Some(ChatDiagnostic::ModelUsageFinish);
}
if restart && usage.restart(durable).await.is_err() {
return Some(ChatDiagnostic::ModelUsageRestart);
}
None
}