kcode-k1-chat-thread-session-inference-settlement 0.1.2

One-shot provider inference settlement for the K1 chat-thread session actor
Documentation
#![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
}