use crate::core::events::TurnRoute;
use crate::models::Usage;
use crate::snapshot::SnapshotRepo;
use std::path::Path;
use std::time::{Duration, Instant};
#[derive(Debug)]
pub struct TurnContext {
pub id: String,
#[allow(dead_code)]
pub started_at: Instant,
pub step: u32,
pub max_steps: u32,
#[allow(dead_code)]
pub cancelled: bool,
pub usage: Usage,
pub(crate) latest_parent_input_tokens: Option<u32>,
pub(crate) pending_route: Option<TurnRoute>,
}
impl TurnContext {
pub fn new(max_steps: u32) -> Self {
Self {
id: uuid::Uuid::new_v4().to_string(),
started_at: Instant::now(),
step: 0,
max_steps,
cancelled: false,
usage: Usage {
input_tokens: 0,
output_tokens: 0,
..Usage::default()
},
latest_parent_input_tokens: None,
pending_route: None,
}
}
pub fn next_step(&mut self) -> bool {
self.step += 1;
self.step <= self.max_steps
}
pub fn at_max_steps(&self) -> bool {
self.step >= self.max_steps
}
#[allow(dead_code)]
pub fn cancel(&mut self) {
self.cancelled = true;
}
#[allow(dead_code)]
pub fn elapsed(&self) -> Duration {
self.started_at.elapsed()
}
pub fn add_usage(&mut self, usage: &Usage) {
self.usage.input_tokens = self.usage.input_tokens.saturating_add(usage.input_tokens);
self.usage.output_tokens = self.usage.output_tokens.saturating_add(usage.output_tokens);
self.usage.prompt_cache_hit_tokens = add_optional_usage(
self.usage.prompt_cache_hit_tokens,
usage.prompt_cache_hit_tokens,
);
self.usage.prompt_cache_miss_tokens = add_optional_usage(
self.usage.prompt_cache_miss_tokens,
usage.prompt_cache_miss_tokens,
);
self.usage.prompt_cache_write_tokens = add_optional_usage(
self.usage.prompt_cache_write_tokens,
usage.prompt_cache_write_tokens,
);
self.usage.reasoning_tokens =
add_optional_usage(self.usage.reasoning_tokens, usage.reasoning_tokens);
self.usage.reasoning_replay_tokens = add_optional_usage(
self.usage.reasoning_replay_tokens,
usage.reasoning_replay_tokens,
);
if let Some(delta) = usage.server_tool_use.as_ref() {
let total = self.usage.server_tool_use.get_or_insert_default();
total.code_execution_requests =
add_optional_usage(total.code_execution_requests, delta.code_execution_requests);
total.tool_search_requests =
add_optional_usage(total.tool_search_requests, delta.tool_search_requests);
}
}
pub fn add_parent_usage(&mut self, usage: &Usage) {
self.latest_parent_input_tokens = (usage.input_tokens > 0).then_some(usage.input_tokens);
self.add_usage(usage);
}
}
fn add_optional_usage(total: Option<u32>, delta: Option<u32>) -> Option<u32> {
match (total, delta) {
(Some(total), Some(delta)) => Some(total.saturating_add(delta)),
(None, Some(delta)) => Some(delta),
(Some(total), None) => Some(total),
(None, None) => None,
}
}
#[cfg(test)]
mod usage_tests {
use super::*;
use crate::models::ServerToolUsage;
#[test]
fn add_usage_preserves_replay_and_saturates_server_tool_counters() {
let mut turn = TurnContext::new(2);
turn.add_usage(&Usage {
reasoning_replay_tokens: Some(u32::MAX - 1),
server_tool_use: Some(ServerToolUsage {
code_execution_requests: Some(u32::MAX),
tool_search_requests: Some(2),
}),
..Usage::default()
});
turn.add_usage(&Usage {
reasoning_replay_tokens: Some(9),
server_tool_use: Some(ServerToolUsage {
code_execution_requests: Some(1),
tool_search_requests: Some(3),
}),
..Usage::default()
});
assert_eq!(turn.usage.reasoning_replay_tokens, Some(u32::MAX));
let server = turn.usage.server_tool_use.expect("server tool usage");
assert_eq!(server.code_execution_requests, Some(u32::MAX));
assert_eq!(server.tool_search_requests, Some(5));
}
fn below_threshold(messages: &[crate::models::Message], turn: &TurnContext) -> bool {
let config = crate::compaction::CompactionConfig {
enabled: true,
token_threshold: 100_000,
..Default::default()
};
!crate::compaction::compaction_pressure_reached_with_billed(
messages,
None,
&config,
turn.latest_parent_input_tokens.map(u64::from),
)
}
#[test]
fn cumulative_low_context_parent_steps_cannot_trigger_compaction() {
let mut turn = TurnContext::new(4);
turn.add_parent_usage(&Usage {
input_tokens: 60_000,
..Usage::default()
});
turn.add_parent_usage(&Usage {
input_tokens: 70_000,
..Usage::default()
});
assert_eq!(turn.usage.input_tokens, 130_000);
assert_eq!(turn.latest_parent_input_tokens, Some(70_000));
assert!(below_threshold(&[], &turn));
}
#[test]
fn child_usage_cannot_replace_parent_context_pressure() {
let mut turn = TurnContext::new(4);
turn.add_parent_usage(&Usage {
input_tokens: 70_000,
..Usage::default()
});
turn.add_usage(&Usage {
input_tokens: 250_000,
..Usage::default()
});
assert_eq!(turn.usage.input_tokens, 320_000);
assert_eq!(turn.latest_parent_input_tokens, Some(70_000));
assert!(below_threshold(&[], &turn));
}
}
const USER_PROMPT_LABEL_MAX: usize = 100;
pub(crate) fn format_snapshot_label(
prefix: &str,
turn_seq: u64,
user_prompt: Option<&str>,
) -> String {
let base = format!("{prefix}:{turn_seq}");
match user_prompt {
None | Some("") => base,
Some(prompt) => match snapshot_label_prompt_snippet(prompt) {
None => base,
Some(snippet) => format!("{base}: {snippet}"),
},
}
}
pub(crate) fn snapshot_label_prompt_snippet(prompt: &str) -> Option<String> {
if prompt.is_empty() {
return None;
}
let first_line = prompt.lines().next().unwrap_or("");
let truncated: String = first_line.chars().take(USER_PROMPT_LABEL_MAX).collect();
if truncated.chars().count() < first_line.chars().count() {
Some(format!("{truncated}…"))
} else {
Some(truncated)
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct ParsedSnapshotLabel {
pub kind: String,
pub seq: Option<u64>,
pub prompt_snippet: Option<String>,
}
pub(crate) fn parse_snapshot_label(label: &str) -> ParsedSnapshotLabel {
let (head, snippet) = match label.split_once(": ") {
Some((head, rest)) => (head, Some(rest.to_string())),
None => (label, None),
};
match head.split_once(':') {
Some((kind, seq)) => ParsedSnapshotLabel {
kind: kind.to_string(),
seq: seq.parse::<u64>().ok(),
prompt_snippet: snippet,
},
None => ParsedSnapshotLabel {
kind: head.to_string(),
seq: None,
prompt_snippet: snippet,
},
}
}
pub fn pre_turn_snapshot(
workspace: &Path,
turn_seq: u64,
cap_bytes: u64,
user_prompt: Option<&str>,
session_id: Option<&str>,
) -> Option<String> {
snapshot_with_label(
workspace,
&format_snapshot_label("pre-turn", turn_seq, user_prompt),
cap_bytes,
session_id,
)
}
pub fn pre_tool_snapshot(
workspace: &Path,
call_id: &str,
cap_bytes: u64,
session_id: Option<&str>,
) -> Option<String> {
snapshot_with_label(workspace, &format!("tool:{call_id}"), cap_bytes, session_id)
}
pub fn post_turn_snapshot(
workspace: &Path,
turn_seq: u64,
cap_bytes: u64,
user_prompt: Option<&str>,
session_id: Option<&str>,
) -> Option<String> {
snapshot_with_label(
workspace,
&format_snapshot_label("post-turn", turn_seq, user_prompt),
cap_bytes,
session_id,
)
}
fn snapshot_with_label(
workspace: &Path,
label: &str,
cap_bytes: u64,
session_id: Option<&str>,
) -> Option<String> {
match SnapshotRepo::open_or_init_with_cap(workspace, cap_bytes) {
Ok(repo) => {
let id = match repo.snapshot_with_session(label, session_id) {
Ok(id) => Some(id.0),
Err(e) => {
tracing::warn!(target: "snapshot", "snapshot '{label}' failed: {e}");
return None;
}
};
if let Err(e) = repo.prune_keep_last_n(crate::snapshot::DEFAULT_MAX_SNAPSHOTS) {
tracing::warn!(target: "snapshot", "snapshot prune failed: {e}");
}
id
}
Err(e) => {
tracing::warn!(target: "snapshot", "snapshot repo init failed: {e}");
maybe_notify_snapshots_disabled_once(workspace, &e);
None
}
}
}
#[allow(clippy::print_stderr)]
fn maybe_notify_snapshots_disabled_once(workspace: &Path, error: &std::io::Error) {
let message = error.to_string();
if !(message.contains("workspace too large for snapshots")
|| message.contains("workspace snapshots are disabled"))
{
return;
}
use std::collections::HashSet;
use std::sync::{Mutex, OnceLock};
static NOTIFIED: OnceLock<Mutex<HashSet<String>>> = OnceLock::new();
let key = workspace.to_string_lossy().into_owned();
let set = NOTIFIED.get_or_init(|| Mutex::new(HashSet::new()));
let Ok(mut guard) = set.lock() else {
return;
};
if !guard.insert(key) {
return;
}
eprintln!(
"warning: workspace snapshots/undo are OFF for {}
{message}
raise `[snapshots] max_workspace_gb` in config.toml (or set it to 0 to disable the cap) to opt in.",
workspace.display()
);
}
#[cfg(test)]
mod snapshot_label_tests {
use super::*;
#[test]
fn label_writer_and_parser_agree_on_prompt_snippet() {
let prompt = "rename the widget\nsecond line is dropped";
let label = format_snapshot_label("pre-turn", 7, Some(prompt));
assert_eq!(label, "pre-turn:7: rename the widget");
let parsed = parse_snapshot_label(&label);
assert_eq!(parsed.kind, "pre-turn");
assert_eq!(parsed.seq, Some(7));
assert_eq!(
parsed.prompt_snippet.as_deref(),
snapshot_label_prompt_snippet(prompt).as_deref(),
"a reader must recover exactly the snippet the writer embedded"
);
}
#[test]
fn truncated_prompt_round_trips_with_its_ellipsis() {
let prompt = "x".repeat(USER_PROMPT_LABEL_MAX + 25);
let label = format_snapshot_label("post-turn", 2, Some(&prompt));
let parsed = parse_snapshot_label(&label);
let snippet = parsed.prompt_snippet.expect("snippet");
assert!(snippet.ends_with('…'));
assert_eq!(snippet.chars().count(), USER_PROMPT_LABEL_MAX + 1);
assert_eq!(
Some(snippet),
snapshot_label_prompt_snippet(&prompt),
"truncated snippets must also round-trip"
);
}
#[test]
fn labels_without_a_prompt_parse_without_inventing_one() {
let label = format_snapshot_label("pre-turn", 3, None);
assert_eq!(label, "pre-turn:3");
let parsed = parse_snapshot_label(&label);
assert_eq!(parsed.kind, "pre-turn");
assert_eq!(parsed.seq, Some(3));
assert_eq!(parsed.prompt_snippet, None);
}
#[test]
fn tool_labels_carry_a_call_id_not_a_sequence() {
let label = format!("tool:{}", "call_abc123");
let parsed = parse_snapshot_label(&label);
assert_eq!(parsed.kind, "tool");
assert_eq!(parsed.seq, None, "a call id is not a turn sequence");
assert_eq!(parsed.prompt_snippet, None);
}
#[test]
fn unrecognized_labels_are_reported_rather_than_dropped() {
let parsed = parse_snapshot_label("manual checkpoint");
assert_eq!(parsed.kind, "manual checkpoint");
assert_eq!(parsed.seq, None);
assert_eq!(parsed.prompt_snippet, None);
}
#[test]
fn empty_prompt_contributes_no_snippet() {
assert_eq!(snapshot_label_prompt_snippet(""), None);
assert_eq!(format_snapshot_label("pre-turn", 1, Some("")), "pre-turn:1");
}
}