use rho_sdk::model::context::estimate_text_tokens;
use serde_json::Value;
pub(super) const COMPLETED_CALL_STRING_CAP_CHARS: usize = 500;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum TranscriptBudget {
Unbounded,
Tokens(u64),
}
impl TranscriptBudget {
pub(super) fn less(self, tokens: u64) -> Self {
match self {
Self::Unbounded => Self::Unbounded,
Self::Tokens(limit) => Self::Tokens(limit.saturating_sub(tokens)),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
#[error("classifier context over budget: needs ~{estimated_tokens} tokens, limit {limit_tokens}")]
pub(crate) struct TranscriptOverBudget {
pub estimated_tokens: u64,
pub limit_tokens: u64,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) enum Retention {
Required,
Droppable,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub(super) struct TranscriptLine {
pub text: String,
pub retention: Retention,
}
pub(super) fn cap_string_leaves(value: &Value, cap_chars: usize) -> Value {
match value {
Value::String(text) => {
let total = text.chars().count();
if total <= cap_chars {
return value.clone();
}
let kept: String = text.chars().take(cap_chars).collect();
Value::String(format!("{kept}…[+{} chars]", total - cap_chars))
}
Value::Array(items) => Value::Array(
items
.iter()
.map(|item| cap_string_leaves(item, cap_chars))
.collect(),
),
Value::Object(fields) => Value::Object(
fields
.iter()
.map(|(name, item)| (name.clone(), cap_string_leaves(item, cap_chars)))
.collect(),
),
Value::Null | Value::Bool(_) | Value::Number(_) => value.clone(),
}
}
pub(super) fn fit_transcript(
history: Vec<TranscriptLine>,
tail: Vec<String>,
budget: TranscriptBudget,
) -> Result<String, TranscriptOverBudget> {
let line_tokens = |text: &str| estimate_text_tokens(text).saturating_add(1);
let mut total: u64 = history
.iter()
.map(|line| line_tokens(&line.text))
.chain(tail.iter().map(|line| line_tokens(line)))
.sum();
let limit = match budget {
TranscriptBudget::Tokens(limit) if total > limit => limit,
TranscriptBudget::Unbounded | TranscriptBudget::Tokens(_) => {
return Ok(join(history.into_iter().map(|line| line.text), tail));
}
};
let mut keep = vec![true; history.len()];
let mut omitted = 0_usize;
let mut marker_tokens = 0;
for (index, line) in history.iter().enumerate() {
if total.saturating_add(marker_tokens) <= limit {
break;
}
if line.retention == Retention::Droppable {
keep[index] = false;
total -= line_tokens(&line.text);
omitted += 1;
marker_tokens = line_tokens(&omitted_marker(omitted));
}
}
let estimated_tokens = total.saturating_add(marker_tokens);
if estimated_tokens > limit {
return Err(TranscriptOverBudget {
estimated_tokens,
limit_tokens: limit,
});
}
let kept = history
.into_iter()
.zip(keep)
.filter_map(|(line, keep)| keep.then_some(line.text));
Ok(join(
std::iter::once(omitted_marker(omitted)).chain(kept),
tail,
))
}
fn omitted_marker(count: usize) -> String {
format!("\"omitted_tool_calls\" count={count}")
}
fn join(history: impl Iterator<Item = String>, tail: Vec<String>) -> String {
history.chain(tail).collect::<Vec<_>>().join("\n")
}