use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use misanthropic::prompt::Prompt;
use misanthropic::tool::{self, MethodDef, Notifications, Tool, Use};
pub const CONTEXT_BUFFER_TOKENS: u64 = 16_000;
#[derive(Debug, Clone, Default)]
pub struct ContextGauge(Arc<AtomicU64>);
impl ContextGauge {
pub fn get(&self) -> u64 {
self.0.load(Ordering::Relaxed)
}
pub fn set(&self, tokens: u64) {
self.0.store(tokens, Ordering::Relaxed);
}
pub fn add(&self, tokens: u64) {
self.0.fetch_add(tokens, Ordering::Relaxed);
}
pub fn fits(&self, tokens: u64, window: u64) -> bool {
self.get() + tokens + CONTEXT_BUFFER_TOKENS <= window
}
pub fn record(&self, usage: &misanthropic::response::Usage) {
let input = usage.input_tokens
+ usage.cache_read_input_tokens.unwrap_or(0)
+ usage.cache_creation_input_tokens.unwrap_or(0);
self.set(input + usage.output_tokens);
}
}
pub(super) fn estimate_tokens(text: &str) -> u64 {
text.len() as u64 / 3
}
pub struct Gauged<T> {
inner: T,
context: ContextGauge,
window: u64,
}
impl<T> Gauged<T> {
pub fn new(inner: T, context: ContextGauge, window: u64) -> Self {
Self {
inner,
context,
window,
}
}
fn admit(&self, call: &str, mut result: tool::Result) -> tool::Result {
let text = result.content.to_string();
let tokens = estimate_tokens(&text);
if self.context.fits(tokens, self.window) {
self.context.add(tokens);
return result;
}
let held = self.context.get();
tracing::warn!(
tool = call,
tokens,
held,
window = self.window,
"tool result would not fit in context; replaced by a note"
);
let note = format!(
"This result is about {tokens} tokens ({} KB); with about {held} \
already in your context it would not fit in your {} token \
window, so it was left out. {}",
text.len() / 1024,
self.window,
smaller(call),
);
self.context.add(estimate_tokens(¬e));
result.content = note.into();
result.is_error = true;
result
}
}
fn smaller(call: &str) -> &'static str {
match call {
"get_content" => {
"Read a post with summary=true for it without its comments, or \
read one comment by its id."
}
"search" | "get_feed" | "get_governance_log" | "get_proposals" => {
"Ask for fewer with a smaller limit."
}
_ => "Ask for something smaller.",
}
}
#[async_trait::async_trait]
impl<T: Tool> Tool for Gauged<T> {
fn name(&self) -> &str {
self.inner.name()
}
fn definitions(&self) -> Vec<MethodDef> {
self.inner.definitions()
}
async fn call(&mut self, call: Use) -> tool::Result {
let name = call.name.clone();
let result = self.inner.call(call).await;
self.admit(&name, result)
}
async fn save_json(&mut self) -> serde_json::Value {
self.inner.save_json().await
}
async fn load_json(
&mut self,
json: serde_json::Value,
) -> Result<(), String> {
self.inner.load_json(json).await
}
fn connect(&mut self, mailbox: tool::Mailbox) {
self.inner.connect(mailbox)
}
async fn on_init(
&mut self,
prompt: &mut Prompt,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
self.inner.on_init(prompt).await
}
async fn on_turn(
&mut self,
prompt: &mut Prompt,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
self.inner.on_turn(prompt).await
}
async fn on_teardown(
&mut self,
prompt: &mut Prompt,
) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
self.inner.on_teardown(prompt).await
}
fn subscribe(&mut self) -> Option<Notifications> {
self.inner.subscribe()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_gauge_counts_cached_input_and_output() {
let gauge = ContextGauge::default();
let usage: misanthropic::response::Usage = serde_json::from_str(
r#"{
"input_tokens": 10,
"cache_read_input_tokens": 90000,
"cache_creation_input_tokens": 1000,
"output_tokens": 500
}"#,
)
.unwrap();
gauge.record(&usage);
assert_eq!(gauge.get(), 91_510);
gauge.add(10);
assert_eq!(gauge.clone().get(), 91_520, "shared, not copied");
}
struct Echo {
len: usize,
}
#[async_trait::async_trait]
impl Tool for Echo {
fn name(&self) -> &str {
"echo"
}
fn definitions(&self) -> Vec<MethodDef> {
Vec::new()
}
async fn call(&mut self, call: Use) -> tool::Result {
tool::Result::new(call.id, "x".repeat(self.len))
}
}
fn call(name: &str) -> Use {
serde_json::from_str(&format!(
r#"{{"id": "toolu_1", "name": "{name}", "input": {{}}}}"#
))
.unwrap()
}
#[tokio::test]
async fn every_result_is_counted_and_an_oversized_one_left_out() {
let gauge = ContextGauge::default();
gauge.set(10_000);
let mut tool = Gauged::new(Echo { len: 30_000 }, gauge.clone(), 40_000);
let first = tool.call(call("search")).await;
assert!(!first.is_error);
assert_eq!(first.content.to_string().len(), 30_000);
assert_eq!(gauge.get(), 20_000);
let second = tool.call(call("search")).await;
assert!(second.is_error);
assert_eq!(second.tool_use_id, "toolu_1", "still answers the call");
let note = second.content.to_string();
assert!(note.contains("about 10000 tokens (29 KB)"), "{note}");
assert!(note.contains("about 20000 already"), "{note}");
assert!(note.contains("smaller limit"), "{note}");
assert!(gauge.get() > 20_000 && gauge.get() < 20_200, "the note");
}
}