diff --git a/.gitignore b/.gitignore index 84a267d..fd41402 100644 --- a/.gitignore +++ b/.gitignore @@ -1,3 +1,4 @@ /target /.direnv /agent.db +/.codex diff --git a/klbr-core/src/agent.rs b/klbr-core/src/agent.rs index 45a00a3..bcfb2cc 100644 --- a/klbr-core/src/agent.rs +++ b/klbr-core/src/agent.rs @@ -59,20 +59,29 @@ pub async fn run( Interrupt::UserMessage(ref text) => { let source = interrupt.source_tag().to_string(); - let memories: Vec = llm - .embed(text) - .await - .ok() - .and_then(|emb| { - memory - .recall(Some(&emb), &[], false, config.memory_top_k) - .ok() - }) - .unwrap_or_default() - .into_iter() - .filter(|e| e.distance.unwrap_or(f32::MAX) < config.memory_sim_threshold) - .map(|e| e.content) - .collect(); + let memories: Vec = match llm.embed(text).await { + Ok(emb) => match memory.recall(Some(&emb), &[], false, config.memory_top_k) { + Ok(results) => results + .into_iter() + .filter(|e| { + e.distance.unwrap_or(f32::MAX) < config.memory_sim_threshold + }) + .map(|e| e.content) + .collect(), + Err(e) => { + let msg = format!("memory recall failed: {e}"); + tracing::error!(%msg); + let _ = output.send(AgentEvent::Error(msg)); + vec![] + } + }, + Err(e) => { + let msg = format!("query embedding failed: {e}"); + tracing::error!(%msg); + let _ = output.send(AgentEvent::Error(msg)); + vec![] + } + }; if !memories.is_empty() { let _ = output.send(AgentEvent::Status(format!( @@ -101,13 +110,12 @@ pub async fn run( let llm2 = llm.clone(); let msgs = ctx.as_messages(); let defs = registry.definitions(); - tokio::spawn(async move { - let _ = llm2.stream(&msgs, &defs, tok_tx).await; - }); + let stream_task = tokio::spawn(async move { llm2.stream(&msgs, &defs, tok_tx).await }); let mut response = String::new(); let mut thinking = String::new(); let mut tool_calls = vec![]; + let mut stream_error = None::; while let Some(ev) = tok_rx.recv().await { match ev { @@ -128,6 +136,23 @@ pub async fn run( } } } + match stream_task.await { + Ok(Ok(())) => {} + Ok(Err(e)) => { + let msg = format!("llm stream failed: {e}"); + tracing::error!(%msg); + let _ = output.send(AgentEvent::Error(msg.clone())); + let _ = output.send(AgentEvent::Status("llm stream failed".into())); + stream_error = Some(msg); + } + Err(e) => { + let msg = format!("llm stream task panicked/cancelled: {e}"); + tracing::error!(%msg); + let _ = output.send(AgentEvent::Error(msg.clone())); + let _ = output.send(AgentEvent::Status("llm stream task failed".into())); + stream_error = Some(msg); + } + } if !tool_calls.is_empty() && tool_iterations < MAX_TOOL_ITERATIONS { tool_iterations += 1; @@ -158,6 +183,11 @@ pub async fn run( continue; } + if stream_error.is_some() && response.is_empty() && tool_calls.is_empty() { + let _ = output.send(AgentEvent::Done); + break; + } + // plain text response (or tool limit hit) — wrap up the turn ctx.push_assistant(&response); let thinking_ref = @@ -314,11 +344,11 @@ async fn reflect( let llm2 = tool_ctx.llm.clone(); let msgs_snap = msgs.clone(); let defs_snap = reflect_registry.definitions(); - tokio::spawn(async move { - let _ = llm2.stream(&msgs_snap, &defs_snap, tok_tx).await; - }); + let stream_task = + tokio::spawn(async move { llm2.stream(&msgs_snap, &defs_snap, tok_tx).await }); let mut tool_calls = vec![]; + let mut stream_failed = false; while let Some(ev) = tok_rx.recv().await { match ev { LlmEvent::ToolCalls(calls) => tool_calls = calls, @@ -331,8 +361,23 @@ async fn reflect( _ => {} } } + match stream_task.await { + Ok(Ok(())) => {} + Ok(Err(e)) => { + let msg = format!("reflection llm stream failed: {e}"); + tracing::error!(%msg); + let _ = output.send(AgentEvent::Error(msg)); + stream_failed = true; + } + Err(e) => { + let msg = format!("reflection llm stream task panicked/cancelled: {e}"); + tracing::error!(%msg); + let _ = output.send(AgentEvent::Error(msg)); + stream_failed = true; + } + } - if tool_calls.is_empty() { + if tool_calls.is_empty() || stream_failed { break; } diff --git a/klbr-core/src/config.rs b/klbr-core/src/config.rs index b3ba5da..4ec7b7b 100644 --- a/klbr-core/src/config.rs +++ b/klbr-core/src/config.rs @@ -2,9 +2,11 @@ pub struct Config { /// llama-server base url pub llm_url: String, - /// embedding server base url + /// embedding server base url (typically a separate process from chat model serving) pub embed_url: String, + /// chat/completions model id. set to "auto" to pick from `GET /v1/models`. pub llm_model: String, + /// embeddings model id. keep this explicit to avoid accidental model/dimension drift. pub embed_model: String, pub watermark_tokens: usize, /// how many recent turns to preserve during compaction @@ -12,7 +14,7 @@ pub struct Config { /// memories to inject per turn pub memory_top_k: usize, /// cosine distance cutoff — only inject memories below this (0=identical, 2=opposite). - /// 0.3 ≈ cosine similarity ≥ 0.7, a reasonable bar for nomic-embed. + /// 0.3 ≈ cosine similarity ≥ 0.7, a reasonable default for bge-m3. pub memory_sim_threshold: f32, pub db_path: String, pub anchor: String, @@ -51,17 +53,17 @@ you can and should combine tags: `["person:mayer", "preference"]` for a preferen impl Default for Config { fn default() -> Self { Self { - llm_url: "http://localhost:1234".into(), - embed_url: "http://localhost:1234".into(), - llm_model: "google/gemma-4-26b-a4b".into(), - embed_model: "nomic-embed-text-v1.5".into(), + llm_url: "http://localhost:8001".into(), + embed_url: "http://localhost:8002".into(), + llm_model: "auto".into(), + embed_model: "auto".into(), watermark_tokens: 32_000, compaction_keep: 10, memory_top_k: 3, memory_sim_threshold: 0.3, db_path: "agent.db".into(), anchor: ANCHOR.into(), - embed_dim: 768, + embed_dim: 1024, } } } diff --git a/klbr-core/src/lib.rs b/klbr-core/src/lib.rs index 8ba818d..c990844 100644 --- a/klbr-core/src/lib.rs +++ b/klbr-core/src/lib.rs @@ -28,6 +28,7 @@ pub enum AgentEvent { ReflectStarted, /// reflection loop finished ReflectDone, + Error(String), Status(String), Metrics(AgentMetrics), ToolCall { diff --git a/klbr-core/src/llm.rs b/klbr-core/src/llm.rs index 78be491..d0862a1 100644 --- a/klbr-core/src/llm.rs +++ b/klbr-core/src/llm.rs @@ -157,14 +157,105 @@ impl LlmClient { } } + fn is_auto_model(name: &str) -> bool { + let normalized = name.trim(); + normalized.is_empty() || normalized.eq_ignore_ascii_case("auto") + } + + fn endpoint(base_url: &str, path: &str) -> String { + format!("{}/{}", base_url.trim_end_matches('/'), path) + } + + async fn fetch_models(&self, base_url: &str) -> Result> { + let url = Self::endpoint(base_url, "v1/models"); + let v = self + .client + .get(url) + .send() + .await? + .error_for_status()? + .json::() + .await?; + + let ids = v["data"] + .as_array() + .map(|arr| { + arr.iter() + .filter_map(|m| m["id"].as_str().map(ToString::to_string)) + .collect::>() + }) + .unwrap_or_default(); + Ok(ids) + } + + fn looks_like_embedding_model(id: &str) -> bool { + let n = id.to_ascii_lowercase(); + [ + "embed", + "embedding", + "text-embedding", + "nomic", + "bge", + "e5", + "gte", + "mxbai", + "minilm", + "jina-emb", + ] + .iter() + .any(|needle| n.contains(needle)) + } + + fn choose_auto_model(candidates: &[String], for_embedding: bool) -> Option { + if candidates.is_empty() { + return None; + } + + if for_embedding { + if let Some(id) = candidates + .iter() + .find(|id| Self::looks_like_embedding_model(id)) + { + return Some(id.clone()); + } + } else if let Some(id) = candidates + .iter() + .find(|id| !Self::looks_like_embedding_model(id)) + { + return Some(id.clone()); + } + + Some(candidates[0].clone()) + } + + async fn resolve_model( + &self, + base_url: &str, + configured: &str, + for_embedding: bool, + ) -> Result { + if !Self::is_auto_model(configured) { + return Ok(configured.to_string()); + } + + let models = self.fetch_models(base_url).await?; + let selected = Self::choose_auto_model(&models, for_embedding) + .ok_or_else(|| anyhow::anyhow!("no models returned from {base_url}/v1/models"))?; + Ok(selected) + } + pub async fn stream( &self, messages: &[Message], tools: &[ToolDef], tok_tx: mpsc::Sender, ) -> Result<()> { + let chat_model = self + .resolve_model(&self.config.llm_url, &self.config.llm_model, false) + .await?; + let mut body = json!({ - "model": self.config.llm_model, + "model": chat_model, "messages": messages, "stream": true, "stream_options": { "include_usage": true } @@ -176,7 +267,7 @@ impl LlmClient { let mut res = self .client - .post(format!("{}/v1/chat/completions", self.config.llm_url)) + .post(Self::endpoint(&self.config.llm_url, "v1/chat/completions")) .json(&body) .send() .await? @@ -280,14 +371,18 @@ impl LlmClient { /// non-streaming completion, used for compaction summaries (no tools) pub async fn complete(&self, messages: &[Message]) -> Result<(String, Usage)> { + let chat_model = self + .resolve_model(&self.config.llm_url, &self.config.llm_model, false) + .await?; + let body = json!({ - "model": self.config.llm_model, + "model": chat_model, "messages": messages, "stream": false, }); let v = self .client - .post(format!("{}/v1/chat/completions", self.config.llm_url)) + .post(Self::endpoint(&self.config.llm_url, "v1/chat/completions")) .json(&body) .send() .await? @@ -302,20 +397,35 @@ impl LlmClient { } pub async fn embed(&self, text: &str) -> Result> { - let body = json!({ "model": self.config.embed_model, "input": text }); + let embed_model = self + .resolve_model(&self.config.embed_url, &self.config.embed_model, true) + .await?; + + let body = json!({ "model": embed_model, "input": text }); let v = self .client - .post(format!("{}/v1/embeddings", self.config.embed_url)) + .post(Self::endpoint(&self.config.embed_url, "v1/embeddings")) .json(&body) .send() .await? .json::() .await?; - v["data"][0]["embedding"] + let embedding: Vec = v["data"][0]["embedding"] .as_array() .ok_or_else(|| anyhow::anyhow!("no embedding in response"))? .iter() .map(|x| Ok(x.as_f64().unwrap_or(0.0) as f32)) - .collect() + .collect::>>()?; + + if embedding.len() != self.config.embed_dim { + return Err(anyhow::anyhow!( + "embedding dimension mismatch: got {}, expected {} (model: {})", + embedding.len(), + self.config.embed_dim, + embed_model + )); + } + + Ok(embedding) } } diff --git a/klbr-daemon/src/daemon.rs b/klbr-daemon/src/daemon.rs index 3c989bb..cd7c4cc 100644 --- a/klbr-daemon/src/daemon.rs +++ b/klbr-daemon/src/daemon.rs @@ -155,6 +155,10 @@ async fn handle( AgentEvent::Done => ServerMsg::Done, AgentEvent::ReflectStarted => ServerMsg::ReflectStarted, AgentEvent::ReflectDone => ServerMsg::ReflectDone, + AgentEvent::Error(content) => { + tracing::error!(%content, "agent error"); + ServerMsg::Error { content } + } AgentEvent::Status(content) => ServerMsg::Status { content }, AgentEvent::Metrics(m) => ServerMsg::Metrics { turn_count: m.turn_count, diff --git a/klbr-ipc/src/lib.rs b/klbr-ipc/src/lib.rs index 23ef4e2..d200508 100644 --- a/klbr-ipc/src/lib.rs +++ b/klbr-ipc/src/lib.rs @@ -36,6 +36,9 @@ pub enum ServerMsg { Done, ReflectStarted, ReflectDone, + Error { + content: String, + }, Status { content: String, }, diff --git a/klbr-tui/src/main.rs b/klbr-tui/src/main.rs index 51137cb..651def1 100644 --- a/klbr-tui/src/main.rs +++ b/klbr-tui/src/main.rs @@ -826,9 +826,17 @@ fn handle_message(app: &mut App, line: String) -> Result<()> { *step = AssistantStep::Done; } } - app.status.clear(); + if !app.status.starts_with("error:") { + app.status.clear(); + } app.stream_start = None; } + ServerMsg::Error { content } => { + app.status = format!("error: {content}"); + app.history + .push(ChatMsg::system(format!("error: {content}"))); + app.snap_to_bottom(); + } ServerMsg::Status { content } => { app.status = content; }