diff --git a/klbr-core/src/agent.rs b/klbr-core/src/agent.rs index 05c3d68..ea8cbb8 100644 --- a/klbr-core/src/agent.rs +++ b/klbr-core/src/agent.rs @@ -8,7 +8,7 @@ use crate::{ config::Config, context::Context, interrupt::Interrupt, - llm::{LlmClient, Message}, + llm::{LlmClient, LlmEvent, Message}, memory::MemoryStore, }; use klbr_ipc::ServerMsg; @@ -69,13 +69,19 @@ pub async fn run( let mut response = String::new(); let mut thinking = String::new(); - while let Some((is_think, tok)) = tok_rx.recv().await { - if is_think { - thinking.push_str(&tok); - let _ = output.send(ServerMsg::ThinkToken { content: tok }); - } else { - response.push_str(&tok); - let _ = output.send(ServerMsg::Token { content: tok }); + while let Some(ev) = tok_rx.recv().await { + match ev { + LlmEvent::ThinkToken(tok) => { + thinking.push_str(&tok); + let _ = output.send(ServerMsg::ThinkToken { content: tok }); + } + LlmEvent::Token(tok) => { + response.push_str(&tok); + let _ = output.send(ServerMsg::Token { content: tok }); + } + LlmEvent::Usage(usage) => { + ctx.update_tokens(usage.total_tokens); + } } } @@ -87,13 +93,13 @@ pub async fn run( let _ = output.send(ServerMsg::Done); let metrics = ServerMsg::Metrics { turn_count, - context_chars: ctx.char_count(), - watermark: config.watermark_chars, + context_tokens: ctx.total_tokens, + watermark: config.watermark_tokens, }; *snapshot.write().await = Some(metrics.clone()); let _ = output.send(metrics); - if ctx.char_count() > config.watermark_chars { + if ctx.total_tokens > config.watermark_tokens { let _ = output.send(ServerMsg::Status { content: "compacting...".into(), }); @@ -122,7 +128,7 @@ async fn compact(config: &Config, llm: &LlmClient, memory: &MemoryStore, ctx: &m "summarize these conversation turns concisely, preserving key facts and topics:\n\n{turns_text}" ))]; - let Ok(summary) = llm.complete(&prompt).await else { + let Ok((summary, _)) = llm.complete(&prompt).await else { return; }; let Ok(emb) = llm.embed(&summary).await else { diff --git a/klbr-core/src/config.rs b/klbr-core/src/config.rs index 1cbcb35..eb553f4 100644 --- a/klbr-core/src/config.rs +++ b/klbr-core/src/config.rs @@ -6,8 +6,7 @@ pub struct Config { pub embed_url: String, pub llm_model: String, pub embed_model: String, - /// char count before context compaction fires - pub watermark_chars: usize, + pub watermark_tokens: usize, /// how many recent turns to preserve during compaction pub compaction_keep: usize, /// memories to inject per turn @@ -32,7 +31,7 @@ impl Default for Config { embed_url: "http://localhost:1234".into(), llm_model: "google/gemma-4-26b-a4b".into(), embed_model: "nomic-embed-text-v1.5".into(), - watermark_chars: 48_000, + watermark_tokens: 32_000, compaction_keep: 10, memory_top_k: 3, memory_sim_threshold: 0.3, diff --git a/klbr-core/src/context.rs b/klbr-core/src/context.rs index 62f6cc8..6c73439 100644 --- a/klbr-core/src/context.rs +++ b/klbr-core/src/context.rs @@ -5,6 +5,7 @@ pub struct Context { anchor: Vec, /// rolling conversation turns pub turns: Vec, + pub total_tokens: usize, } impl Context { @@ -12,6 +13,7 @@ impl Context { Self { anchor: vec![Message::system(anchor)], turns: vec![], + total_tokens: 0, } } @@ -46,10 +48,6 @@ impl Context { self.anchor.iter().chain(&self.turns).cloned().collect() } - pub fn char_count(&self) -> usize { - self.turns.iter().map(|m| m.content.len()).sum() - } - pub fn turn_count(&self) -> usize { self.turns.len() } @@ -59,6 +57,13 @@ impl Context { if self.turns.len() <= keep { return vec![]; } + // we reset tokens because we lost track of the exact window count after draining + // it will be updated on the next API call anyway + self.total_tokens = 0; self.turns.drain(..self.turns.len() - keep).collect() } + + pub fn update_tokens(&mut self, tokens: usize) { + self.total_tokens = tokens; + } } diff --git a/klbr-core/src/llm.rs b/klbr-core/src/llm.rs index d971f10..c74b817 100644 --- a/klbr-core/src/llm.rs +++ b/klbr-core/src/llm.rs @@ -34,8 +34,19 @@ impl Message { } } -/// (is_think, token) - is_think=true means reasoning_content (thinking trace) -pub type StreamChunk = (bool, String); +#[derive(Debug, Serialize, Deserialize, Clone, Default)] +pub struct Usage { + pub prompt_tokens: usize, + pub completion_tokens: usize, + pub total_tokens: usize, +} + +#[derive(Debug, Clone)] +pub enum LlmEvent { + Token(String), + ThinkToken(String), + Usage(Usage), +} #[derive(Clone)] pub struct LlmClient { @@ -51,15 +62,12 @@ impl LlmClient { } } - pub async fn stream( - &self, - messages: &[Message], - tok_tx: mpsc::Sender, - ) -> Result<()> { + pub async fn stream(&self, messages: &[Message], tok_tx: mpsc::Sender) -> Result<()> { let body = json!({ "model": self.config.llm_model, "messages": messages, "stream": true, + "stream_options": { "include_usage": true } }); let mut res = self @@ -88,16 +96,34 @@ impl LlmClient { let Ok(v) = serde_json::from_str::(data) else { continue; }; - let delta = &v["choices"][0]["delta"]; + + if let Some(_usage) = v["usage"].as_object() { + let u: Usage = serde_json::from_value(v["usage"].clone())?; + let _ = tok_tx.send(LlmEvent::Usage(u)).await; + continue; + } + + let Some(choices) = v["choices"].as_array() else { + continue; + }; + if choices.is_empty() { + continue; + } + let delta = &choices[0]["delta"]; // reasoning_content comes before content during thinking if let Some(t) = delta["reasoning_content"].as_str() { - if !t.is_empty() && tok_tx.send((true, t.to_string())).await.is_err() { + if !t.is_empty() + && tok_tx + .send(LlmEvent::ThinkToken(t.to_string())) + .await + .is_err() + { return Ok(()); } } if let Some(t) = delta["content"].as_str() { - if !t.is_empty() && tok_tx.send((false, t.to_string())).await.is_err() { + if !t.is_empty() && tok_tx.send(LlmEvent::Token(t.to_string())).await.is_err() { return Ok(()); } } @@ -108,7 +134,7 @@ impl LlmClient { } /// non-streaming, used for compaction summaries - pub async fn complete(&self, messages: &[Message]) -> Result { + pub async fn complete(&self, messages: &[Message]) -> Result<(String, Usage)> { let body = json!({ "model": self.config.llm_model, "messages": messages, @@ -122,10 +148,12 @@ impl LlmClient { .await? .json::() .await?; - Ok(v["choices"][0]["message"]["content"] + let content = v["choices"][0]["message"]["content"] .as_str() .unwrap_or("") - .to_string()) + .to_string(); + let usage: Usage = serde_json::from_value(v["usage"].clone()).unwrap_or_default(); + Ok((content, usage)) } pub async fn embed(&self, text: &str) -> Result> { diff --git a/klbr-ipc/src/lib.rs b/klbr-ipc/src/lib.rs index 6ca55ce..97fd19c 100644 --- a/klbr-ipc/src/lib.rs +++ b/klbr-ipc/src/lib.rs @@ -26,7 +26,7 @@ pub enum ServerMsg { /// sent after each Done with current agent stats Metrics { turn_count: usize, - context_chars: usize, + context_tokens: usize, watermark: usize, }, /// pushed on connect: current context window turns diff --git a/klbr-tui/src/main.rs b/klbr-tui/src/main.rs index 66c6725..8b4047f 100644 --- a/klbr-tui/src/main.rs +++ b/klbr-tui/src/main.rs @@ -111,7 +111,7 @@ struct App { // metrics turn_count: usize, - context_chars: usize, + context_tokens: usize, watermark: usize, stream_start: Option, stream_tokens: usize, @@ -132,7 +132,7 @@ impl App { history_exhausted: false, loading_history: false, turn_count: 0, - context_chars: 0, + context_tokens: 0, watermark: 0, stream_start: None, stream_tokens: 0, @@ -471,13 +471,12 @@ pub async fn run() -> Result<()> { .unwrap_or_default(); let ctx_pct = (app.watermark > 0) - .then(|| (app.context_chars as f64 / app.watermark as f64 * 100.0) as usize) + .then(|| (app.context_tokens as f64 / app.watermark as f64 * 100.0) as usize) .unwrap_or(0); let context_str = if app.watermark > 0 { - let remaining = app.watermark.saturating_sub(app.context_chars); - let tokens_left = remaining / 4; - format!("ctx {ctx_pct}% (~{tokens_left} tok until compact)") + let remaining = app.watermark.saturating_sub(app.context_tokens); + format!("ctx {ctx_pct}% ({remaining} tok left)") } else { String::new() }; @@ -638,11 +637,11 @@ fn handle_message(app: &mut App, line: String) -> Result<()> { } ServerMsg::Metrics { turn_count, - context_chars, + context_tokens, watermark, } => { app.turn_count = turn_count; - app.context_chars = context_chars; + app.context_tokens = context_tokens; app.watermark = watermark; } ServerMsg::History { turns } => {