use std::collections::BTreeMap;
use std::sync::Arc;
use std::time::Instant;
use tiktoken_rs::CoreBPE;
use tracing::{debug, error};
use scryer_db::{ScryerDb, ToolInvocationMetric};
pub const MAX_PAYLOAD_TOKENS: usize = 1000;
pub const COST_PER_MILLION_TOKENS: f64 = 2.50;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ToolCategory {
OutlineOrScope,
DefinitionOrType,
ReferencesOrHierarchy,
InspectSymbol,
CrateOutline,
SearchSymbols,
Default,
}
impl ToolCategory {
pub fn for_tool(tool_name: &str) -> Self {
match tool_name {
"get_file_outline" | "get_enclosing_scope" => Self::OutlineOrScope,
"resolve_definition" | "get_type_contract" => Self::DefinitionOrType,
"find_references" | "trace_call_hierarchy" | "calculate_blast_radius" => {
Self::ReferencesOrHierarchy
}
"inspect_symbol" => Self::InspectSymbol,
"get_crate_outline" => Self::CrateOutline,
"search_symbols" => Self::SearchSymbols,
_ => Self::Default,
}
}
}
#[derive(Clone)]
pub struct TokenSavingsMiddleware {
db: ScryerDb,
bpe: Arc<CoreBPE>,
}
impl TokenSavingsMiddleware {
pub fn new(db: ScryerDb) -> anyhow::Result<Self> {
let bpe = tiktoken_rs::cl100k_base()?;
Ok(Self {
db,
bpe: Arc::new(bpe),
})
}
pub fn bpe(&self) -> &CoreBPE {
&self.bpe
}
pub fn count_tokens(&self, text: &str) -> usize {
self.bpe.encode_with_special_tokens(text).len()
}
pub fn enforce_limit(
&self,
payload: &str,
max_tokens: Option<usize>,
no_truncate: bool,
) -> (String, usize, bool) {
let count = self.count_tokens(payload);
if no_truncate {
return (payload.to_string(), count, false);
}
let budget = max_tokens.unwrap_or(MAX_PAYLOAD_TOKENS);
if budget == 0 || count <= budget {
return (payload.to_string(), count, false);
}
let mut truncated = String::new();
let budget_str = if budget == 1000 {
"1,000".to_string()
} else {
budget.to_string()
};
let notice = format!(
"\n\n... [Output truncated to enforce hard {budget_str} token limit. Pass 'no_truncate: true', increase 'max_tokens', or use pagination/filtering (offset, limit, query).]"
);
let notice_tokens = self.count_tokens(¬ice);
let token_budget = budget.saturating_sub(notice_tokens);
for line in payload.lines() {
let next = if truncated.is_empty() {
line.to_string()
} else {
format!("{truncated}\n{line}")
};
if self.count_tokens(&next) > token_budget {
break;
}
truncated = next;
}
truncated.push_str(¬ice);
let actual = self.count_tokens(&truncated);
(truncated, actual, true)
}
pub fn enforce_default_limit(&self, payload: &str) -> (String, usize, bool) {
self.enforce_limit(payload, None, false)
}
pub fn calculate_naive_cost(
category: ToolCategory,
file_bytes: Option<usize>,
match_count: Option<usize>,
) -> u64 {
let bytes = file_bytes.unwrap_or(0) as f64;
let naive_bytes_tokens = bytes / 3.8;
match category {
ToolCategory::OutlineOrScope => {
if file_bytes.is_some() {
naive_bytes_tokens.round() as u64
} else {
1000
}
}
ToolCategory::DefinitionOrType => naive_bytes_tokens.max(1500.0).round() as u64,
ToolCategory::ReferencesOrHierarchy => {
let matches = (match_count.unwrap_or(1).min(5)) as f64;
let per_file = naive_bytes_tokens.max(1200.0);
(matches * per_file).round() as u64
}
ToolCategory::InspectSymbol => 6000,
ToolCategory::CrateOutline => 8000,
ToolCategory::SearchSymbols => {
let matches = (match_count.unwrap_or(1).min(5)) as f64;
(matches * 1200.0).max(3500.0).round() as u64
}
ToolCategory::Default => 1000,
}
}
pub fn record_invocation(
&self,
session_id: String,
project_id: Option<u64>,
tool_name: String,
payload_tokens: u64,
estimated_naive_tokens: u64,
duration: Instant,
) {
let duration_ms = duration.elapsed().as_millis() as u32;
let tokens_saved = estimated_naive_tokens.saturating_sub(payload_tokens);
let db = self.db.clone();
let now = scryer_db::time::now_rfc3339();
tokio::spawn(async move {
let mut guard = db.lock().await;
let result = ToolInvocationMetric::create()
.session_id(session_id)
.project_id(project_id)
.tool_name(tool_name)
.payload_tokens(payload_tokens)
.estimated_naive_tokens(estimated_naive_tokens)
.tokens_saved(tokens_saved)
.execution_duration_ms(duration_ms)
.created_at(now)
.exec(&mut *guard)
.await;
if let Err(e) = result {
error!("Failed to persist ToolInvocationMetric: {e}");
} else {
debug!(tokens_saved, "Persisted ToolInvocationMetric successfully");
}
});
}
pub async fn get_metrics(
&self,
session_id: Option<&str>,
) -> anyhow::Result<TokenSavingsSummary> {
let mut guard = self.db.lock().await;
let metrics = match session_id {
Some(sid) => {
ToolInvocationMetric::filter(ToolInvocationMetric::fields().session_id().eq(sid))
.exec(&mut *guard)
.await?
}
None => ToolInvocationMetric::all().exec(&mut *guard).await?,
};
let total_invocations = metrics.len() as u64;
let mut total_payload_tokens = 0u64;
let mut total_naive_tokens = 0u64;
let mut total_tokens_saved = 0u64;
let mut total_duration_ms = 0u64;
let mut by_tool: BTreeMap<String, u64> = BTreeMap::new();
for m in &metrics {
*by_tool.entry(m.tool_name.clone()).or_default() += 1;
total_payload_tokens += m.payload_tokens;
total_naive_tokens += m.estimated_naive_tokens;
total_tokens_saved += m.tokens_saved;
total_duration_ms += m.execution_duration_ms as u64;
}
let average_duration_ms = if total_invocations > 0 {
total_duration_ms as f64 / total_invocations as f64
} else {
0.0
};
let billable_cost_reduction_dollars =
(total_tokens_saved as f64 / 1_000_000.0) * COST_PER_MILLION_TOKENS;
Ok(TokenSavingsSummary {
total_invocations,
total_payload_tokens,
total_estimated_naive_tokens: total_naive_tokens,
total_tokens_saved,
average_duration_ms,
billable_cost_reduction_dollars,
by_tool,
})
}
}
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, schemars::JsonSchema)]
pub struct TokenSavingsSummary {
pub total_invocations: u64,
pub total_payload_tokens: u64,
pub total_estimated_naive_tokens: u64,
pub total_tokens_saved: u64,
pub average_duration_ms: f64,
pub billable_cost_reduction_dollars: f64,
#[serde(default)]
pub by_tool: BTreeMap<String, u64>,
}