use std::collections::BTreeMap;
use serde::Deserialize;
use serde::Serialize;
use serde_json::Value;
use crate::BoxFuture;
use crate::Error;
use crate::Result;
use crate::backend::model::ToolCall;
use crate::backend::sandbox::NetworkAccess;
use crate::protocol::MessageTarget;
use crate::protocol::SessionContext;
use crate::protocol::TokenUsage;
pub mod sqlite;
pub(crate) const CHECKPOINT_VERSION: u32 = 4;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ActiveExecution {
pub submission_id: String,
pub turn_id: String,
pub started_at_ms: i64,
pub model_calls: u64,
pub tool_calls: u64,
pub failed_tool_calls: u64,
pub usage: TokenUsage,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ExecutionOutcome {
Completed,
Aborted,
Failed,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecutionRecord {
pub session_id: String,
pub submission_id: String,
pub turn_id: String,
pub started_at_ms: i64,
pub finished_at_ms: i64,
pub elapsed_ms: u64,
pub outcome: ExecutionOutcome,
pub model_calls: u64,
pub tool_calls: u64,
pub failed_tool_calls: u64,
pub usage: TokenUsage,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecutionStats {
pub run_count: u64,
pub failed_run_count: u64,
pub aborted_run_count: u64,
pub model_calls: u64,
pub tool_calls: u64,
pub failed_tool_calls: u64,
pub elapsed_ms: u64,
pub usage: TokenUsage,
}
impl ExecutionStats {
pub(crate) fn checked_record(&mut self, record: &ExecutionRecord) -> Option<()> {
let run_count = self.run_count.checked_add(1)?;
let failed_run_count = self
.failed_run_count
.checked_add(u64::from(record.outcome == ExecutionOutcome::Failed))?;
let aborted_run_count = self
.aborted_run_count
.checked_add(u64::from(record.outcome == ExecutionOutcome::Aborted))?;
let model_calls = self.model_calls.checked_add(record.model_calls)?;
let tool_calls = self.tool_calls.checked_add(record.tool_calls)?;
let failed_tool_calls = self
.failed_tool_calls
.checked_add(record.failed_tool_calls)?;
let elapsed_ms = self.elapsed_ms.checked_add(record.elapsed_ms)?;
let mut usage = self.usage.clone();
usage.checked_add(&record.usage)?;
*self = Self {
run_count,
failed_run_count,
aborted_run_count,
model_calls,
tool_calls,
failed_tool_calls,
elapsed_ms,
usage,
};
Some(())
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PendingApproval {
pub submission_id: String,
pub turn_id: String,
pub request_id: String,
pub approval_call_ids: Vec<String>,
pub authorized_call_ids: Vec<String>,
pub calls: Vec<ToolCall>,
pub reason: String,
pub network_access: NetworkAccess,
pub decision_received: bool,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct Checkpoint {
pub version: u32,
pub session_id: String,
pub session_context: SessionContext,
pub metadata: BTreeMap<String, Value>,
pub catalog_visible: bool,
pub first_user_message: Option<String>,
pub model_route: Option<String>,
pub sequence: u64,
pub context: Vec<Value>,
pub total_usage: TokenUsage,
pub last_usage: Option<TokenUsage>,
pub pending_input: Vec<String>,
pub active_execution: Option<ActiveExecution>,
pub execution_stats: ExecutionStats,
pub pending_tools: Vec<ToolCall>,
pub pending_approval: Option<PendingApproval>,
}
impl Checkpoint {
#[must_use]
pub fn empty(session_id: impl Into<String>) -> Self {
Self {
version: CHECKPOINT_VERSION,
session_id: session_id.into(),
session_context: SessionContext::default(),
metadata: BTreeMap::new(),
catalog_visible: true,
first_user_message: None,
model_route: None,
sequence: 0,
context: Vec::new(),
total_usage: TokenUsage::default(),
last_usage: None,
pending_input: Vec::new(),
active_execution: None,
execution_stats: ExecutionStats::default(),
pending_tools: Vec::new(),
pending_approval: None,
}
}
pub(crate) fn finish_execution(
&mut self,
outcome: ExecutionOutcome,
finished_at_ms: i64,
) -> Result<ExecutionRecord> {
let active = self
.active_execution
.as_ref()
.ok_or_else(|| Error::Checkpoint("turn ended without an active execution".into()))?;
let finished_at_ms = finished_at_ms.max(active.started_at_ms);
let elapsed_ms = u64::try_from(finished_at_ms - active.started_at_ms)
.map_err(|_| Error::Checkpoint("execution elapsed time is unsupported".into()))?;
let record = ExecutionRecord {
session_id: self.session_id.clone(),
submission_id: active.submission_id.clone(),
turn_id: active.turn_id.clone(),
started_at_ms: active.started_at_ms,
finished_at_ms,
elapsed_ms,
outcome,
model_calls: active.model_calls,
tool_calls: active.tool_calls,
failed_tool_calls: active.failed_tool_calls,
usage: active.usage.clone(),
};
let mut stats = self.execution_stats.clone();
stats.checked_record(&record).ok_or_else(|| {
Error::Checkpoint("execution statistics exceed the supported range".into())
})?;
self.active_execution = None;
self.execution_stats = stats;
Ok(record)
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SessionSummary {
pub session_id: String,
pub session_context: SessionContext,
pub parent_session_id: Option<String>,
pub parent_sequence: Option<u64>,
pub sequence: u64,
pub catalog_visible: bool,
pub first_user_message: Option<String>,
pub execution_stats: ExecutionStats,
pub created_at: i64,
pub updated_at: i64,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SessionCursor {
pub updated_at: i64,
pub sequence: u64,
pub session_id: String,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SessionPageRequest {
pub cursor: Option<SessionCursor>,
pub limit: usize,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct SessionPage {
pub sessions: Vec<SessionSummary>,
pub next_cursor: Option<SessionCursor>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TranscriptBatch {
pub sequence: u64,
pub created_at: i64,
pub items: Vec<Value>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecutionPageRequest {
pub before_sequence: Option<u64>,
pub limit: usize,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExecutionPage {
pub executions: Vec<ExecutionRecord>,
pub next_before_sequence: Option<u64>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct TranscriptPageRequest {
pub before_sequence: Option<u64>,
pub max_batches: usize,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct TranscriptPage {
pub batches: Vec<TranscriptBatch>,
pub next_before_sequence: Option<u64>,
}
impl TranscriptPage {
#[must_use]
pub fn into_positioned_items_chronological(self) -> Vec<(MessageTarget, Value)> {
self.batches
.into_iter()
.rev()
.flat_map(|batch| {
batch
.items
.into_iter()
.enumerate()
.map(move |(index, item)| {
(
MessageTarget {
checkpoint_sequence: batch.sequence,
batch_item_count: index + 1,
},
item,
)
})
})
.collect()
}
}
pub trait CheckpointStore: Send + Sync {
fn load<'a>(&'a self, session_id: &'a str) -> BoxFuture<'a, Result<Option<Checkpoint>>>;
fn save<'a>(
&'a self,
checkpoint: &'a Checkpoint,
transcript_delta: &'a [Value],
execution: Option<&'a ExecutionRecord>,
) -> BoxFuture<'a, Result<()>>;
fn list_sessions_page(
&self,
_request: SessionPageRequest,
) -> BoxFuture<'_, Result<SessionPage>> {
Box::pin(async {
Err(Error::Checkpoint(
"this checkpoint backend has no session catalog".into(),
))
})
}
fn transcript_page<'a>(
&'a self,
session_id: &'a str,
request: TranscriptPageRequest,
) -> BoxFuture<'a, Result<TranscriptPage>> {
Box::pin(async move {
if request.max_batches == 0 {
return Err(Error::Checkpoint(
"transcript page limit must be positive".into(),
));
}
let Some(checkpoint) = self.load(session_id).await? else {
return Ok(TranscriptPage::default());
};
if checkpoint.context.is_empty()
|| request
.before_sequence
.is_some_and(|before| checkpoint.sequence >= before)
{
return Ok(TranscriptPage::default());
}
Ok(TranscriptPage {
batches: vec![TranscriptBatch {
sequence: checkpoint.sequence,
created_at: 0,
items: checkpoint.context,
}],
next_before_sequence: None,
})
})
}
fn execution_page<'a>(
&'a self,
_session_id: &'a str,
_request: ExecutionPageRequest,
) -> BoxFuture<'a, Result<ExecutionPage>> {
Box::pin(async {
Err(Error::Checkpoint(
"this checkpoint backend has no execution journal".into(),
))
})
}
fn recent_executions(&self, _limit: usize) -> BoxFuture<'_, Result<Vec<ExecutionRecord>>> {
Box::pin(async {
Err(Error::Checkpoint(
"this checkpoint backend has no execution journal".into(),
))
})
}
fn fork<'a>(
&'a self,
_parent_session_id: &'a str,
_parent_sequence: u64,
_checkpoint: &'a Checkpoint,
) -> BoxFuture<'a, Result<SessionSummary>> {
Box::pin(async {
Err(Error::Checkpoint(
"this checkpoint backend cannot fork sessions".into(),
))
})
}
fn load_state<'a>(
&'a self,
scope: &'a str,
key: &'a str,
) -> BoxFuture<'a, Result<Option<Value>>>;
fn save_state<'a>(
&'a self,
scope: &'a str,
key: &'a str,
value: &'a Value,
) -> BoxFuture<'a, Result<()>>;
}