From 0450aaa39ff5bcee6bd3057299389a19793e7d2d Mon Sep 17 00:00:00 2001 From: Lifei Zhou Date: Wed, 8 Apr 2026 19:51:57 +1000 Subject: [PATCH] enable gemma 4 local model (#8366) Co-authored-by: Claude Opus 4.6 (1M context) --- Cargo.lock | 12 +- .../src/routes/local_inference.rs | 17 +- crates/goose/Cargo.toml | 4 +- crates/goose/src/providers/local_inference.rs | 19 +- .../local_inference/inference_native_tools.rs | 309 ++++++++-- .../local_inference/local_model_registry.rs | 57 +- .../providers/local_inference/tool_parsing.rs | 546 +----------------- .../models/bottom_bar/ModelsBottomBar.tsx | 2 +- 8 files changed, 333 insertions(+), 633 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 1346958d..9ef30097 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1724,9 +1724,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.2.57" +version = "1.2.59" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7a0dd1ca384932ff3641c8718a02769f1698e7563dc6974ffd03346116310423" +checksum = "b7a4d3ec6524d28a329fc53654bbadc9bdd7b0431f5d65f1a56ffb28a1ee5283" dependencies = [ "find-msvc-tools", "jobserver", @@ -5860,9 +5860,9 @@ checksum = "11d3d7f243d5c5a8b9bb5d6dd2b1602c0cb0b9db1621bafc7ed66e35ff9fe092" [[package]] name = "llama-cpp-2" -version = "0.1.140" +version = "0.1.142" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e5604c13b9c847157470479a64d1d7c94f3089709309f82f2fdcbcd43510f2f2" +checksum = "9f04fe7a4b6836bff1966eb8c35e5e121f573641887526d9f997e9c10a64d1bc" dependencies = [ "encoding_rs", "enumflags2", @@ -5874,9 +5874,9 @@ dependencies = [ [[package]] name = "llama-cpp-sys-2" -version = "0.1.140" +version = "0.1.142" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cbdd3e2c06f3a9a47466a631735946e9ad47fef565b88bc8766a3794474a66f3" +checksum = "95ef4708a34c686a2686f9a57050df3b5384ff346dcbfb8ee47b5814f4dcfc90" dependencies = [ "bindgen", "cc", diff --git a/crates/goose-server/src/routes/local_inference.rs b/crates/goose-server/src/routes/local_inference.rs index d93518e3..ddad43ff 100644 --- a/crates/goose-server/src/routes/local_inference.rs +++ b/crates/goose-server/src/routes/local_inference.rs @@ -13,8 +13,9 @@ use goose::providers::local_inference::{ available_inference_memory_bytes, hf_models::{resolve_model_spec, HfGgufFile}, local_model_registry::{ - get_registry, is_featured_model, model_id_from_repo, LocalModelEntry, - ModelDownloadStatus as RegistryDownloadStatus, ModelSettings, FEATURED_MODELS, + default_settings_for_model, get_registry, is_featured_model, model_id_from_repo, + LocalModelEntry, ModelDownloadStatus as RegistryDownloadStatus, ModelSettings, + FEATURED_MODELS, }, recommend_local_model, }; @@ -51,8 +52,8 @@ pub struct LocalModelResponse { async fn ensure_featured_models_in_registry() -> Result<(), ErrorResponse> { let mut entries_to_add = Vec::new(); - for spec in FEATURED_MODELS { - let (repo_id, quantization) = match hf_models::parse_model_spec(spec) { + for featured in FEATURED_MODELS { + let (repo_id, quantization) = match hf_models::parse_model_spec(featured.spec) { Ok(parts) => parts, Err(_) => continue, }; @@ -68,7 +69,7 @@ async fn ensure_featured_models_in_registry() -> Result<(), ErrorResponse> { } } - let hf_file = match resolve_model_spec(spec).await { + let hf_file = match resolve_model_spec(featured.spec).await { Ok((_repo, file)) => file, Err(_) => { let filename = format!( @@ -91,13 +92,13 @@ async fn ensure_featured_models_in_registry() -> Result<(), ErrorResponse> { let local_path = Paths::in_data_dir("models").join(&hf_file.filename); entries_to_add.push(LocalModelEntry { - id: model_id, + id: model_id.clone(), repo_id, filename: hf_file.filename, quantization, local_path, source_url: hf_file.download_url, - settings: ModelSettings::default(), + settings: default_settings_for_model(&model_id), size_bytes: hf_file.size_bytes, }); } @@ -273,7 +274,7 @@ pub async fn download_hf_model( quantization, local_path: local_path.clone(), source_url: download_url.clone(), - settings: ModelSettings::default(), + settings: default_settings_for_model(&model_id), size_bytes: hf_file.size_bytes, }; diff --git a/crates/goose/Cargo.toml b/crates/goose/Cargo.toml index a3a68b69..0a892e44 100644 --- a/crates/goose/Cargo.toml +++ b/crates/goose/Cargo.toml @@ -179,7 +179,7 @@ tree-sitter-typescript = { workspace = true } which = { workspace = true } pctx_code_mode = { version = "^0.3.0", optional = true } pulldown-cmark = "0.13.0" -llama-cpp-2 = { version = "0.1.137", features = ["sampler"], optional = true } +llama-cpp-2 = { version = "0.1.142", features = ["sampler"], optional = true } encoding_rs = "0.8.35" pastey = "0.2.1" shell-words = { workspace = true } @@ -197,7 +197,7 @@ keyring = { version = "3.6.2", features = ["windows-native"] } [target.'cfg(target_os = "macos")'.dependencies] candle-core = { version = "0.9", default-features = false, features = ["metal"], optional = true } candle-nn = { version = "0.9", default-features = false, features = ["metal"], optional = true } -llama-cpp-2 = { version = "0.1.137", features = ["sampler", "metal"], optional = true } +llama-cpp-2 = { version = "0.1.142", features = ["sampler", "metal"], optional = true } keyring = { version = "3.6.2", features = ["apple-native"] } [target.'cfg(target_os = "linux")'.dependencies] diff --git a/crates/goose/src/providers/local_inference.rs b/crates/goose/src/providers/local_inference.rs index 04bac00f..f6b6fb97 100644 --- a/crates/goose/src/providers/local_inference.rs +++ b/crates/goose/src/providers/local_inference.rs @@ -191,7 +191,7 @@ pub fn recommend_local_model(runtime: &InferenceRuntime) -> String { } // Fallback to first featured model - FEATURED_MODELS[0].to_string() + FEATURED_MODELS[0].spec.to_string() } fn build_openai_messages_json(system: &str, messages: &[Message]) -> String { @@ -360,7 +360,7 @@ impl LocalInferenceProvider { } }; - tracing::info!("Model loaded successfully"); + tracing::info!(model_id = model_id, "Model loaded successfully"); Ok(LoadedModel { model, template }) } @@ -377,7 +377,7 @@ impl ProviderDef for LocalInferenceProvider { get_registry, FEATURED_MODELS, }; - let mut known_models: Vec<&str> = FEATURED_MODELS.to_vec(); + let mut known_models: Vec<&str> = FEATURED_MODELS.iter().map(|m| m.spec).collect(); // Add any registry models not already in the featured list let mut dynamic_models = Vec::new(); @@ -477,13 +477,10 @@ impl Provider for LocalInferenceProvider { } } - // Models that support native OpenAI-compatible tool-call JSON use the - // native path (template-based tool calling with JSON output). All other - // models use the emulator which parses `$ command` and ```execute blocks. - // Only use emulator when there are actually tools to emulate - utility calls - // like compaction and session naming pass empty tools and should preserve - // their system prompts. - let use_emulator = !model_settings.native_tool_calling && !tools.is_empty(); + // Use the model's native_tool_calling setting to decide the path. + // Featured models have this set explicitly; user-added models default to false. + let native_tool_calling = model_settings.native_tool_calling; + let use_emulator = !native_tool_calling && !tools.is_empty(); let system_prompt = if use_emulator { load_tiny_model_prompt() } else { @@ -539,7 +536,7 @@ impl Provider for LocalInferenceProvider { (None, None) }; - let oai_messages_json = if model_settings.use_jinja { + let oai_messages_json = if model_settings.use_jinja || native_tool_calling { Some(build_openai_messages_json(&system_prompt, messages)) } else { None diff --git a/crates/goose/src/providers/local_inference/inference_native_tools.rs b/crates/goose/src/providers/local_inference/inference_native_tools.rs index 656a6e08..1894c4b6 100644 --- a/crates/goose/src/providers/local_inference/inference_native_tools.rs +++ b/crates/goose/src/providers/local_inference/inference_native_tools.rs @@ -1,17 +1,17 @@ -use crate::conversation::message::Message; +use crate::conversation::message::{Message, MessageContent}; use crate::providers::errors::ProviderError; use llama_cpp_2::model::AddBos; use llama_cpp_2::openai::OpenAIChatTemplateParams; +use rmcp::model::CallToolRequestParams; +use serde_json::Value; +use std::borrow::Cow; +use uuid::Uuid; use super::finalize_usage; use super::inference_engine::{ context_cap, create_and_prefill_context, estimate_max_context_for_memory, generation_loop, validate_and_compute_context, GenerationContext, TokenAction, }; -use super::tool_parsing::{ - extract_tool_call_messages, extract_xml_tool_call_messages, safe_stream_end, - split_content_and_tool_calls, split_content_and_xml_tool_calls, -}; pub(super) fn generate_with_native_tools( ctx: &mut GenerationContext<'_>, @@ -105,7 +105,20 @@ pub(super) fn generate_with_native_tools( let message_id = ctx.message_id; let tx = ctx.tx; let mut generated_text = String::new(); - let mut streamed_len: usize = 0; + + // Initialize streaming parser — handles thinking tokens, tool calls, etc. + let mut stream_parser = template_result.streaming_state_oaicompat().map_err(|e| { + ProviderError::ExecutionError(format!("Failed to init streaming parser: {}", e)) + })?; + + // Feed the generation prompt to the parser so it knows the context. + // The model may echo this prefix; the parser needs to see it to strip it. + if !template_result.generation_prompt.is_empty() { + let _ = stream_parser.update(&template_result.generation_prompt, true); + } + + // Accumulate tool calls across streaming deltas + let mut accumulated_tool_calls: Vec = Vec::new(); let output_token_count = generation_loop( &ctx.loaded.model, @@ -116,28 +129,35 @@ pub(super) fn generate_with_native_tools( |piece| { generated_text.push_str(piece); - let has_xml_tc = split_content_and_xml_tool_calls(&generated_text).is_some(); - let (content, tc) = split_content_and_tool_calls(&generated_text); - let stream_up_to = if tc.is_some() { - content.len() - } else if has_xml_tc { - split_content_and_xml_tool_calls(&generated_text) - .map(|(c, _)| c.len()) - .unwrap_or(0) - } else { - safe_stream_end(&generated_text) - }; - if stream_up_to > streamed_len { - #[allow(clippy::string_slice)] - let new_text = &generated_text[streamed_len..stream_up_to]; - if !new_text.is_empty() { - let mut msg = Message::assistant().with_text(new_text); - msg.id = Some(message_id.to_string()); - if tx.blocking_send(Ok((Some(msg), None))).is_err() { - return Ok(TokenAction::Stop); + // Feed the new piece to the streaming parser + match stream_parser.update(piece, true) { + Ok(deltas) => { + for delta_json in deltas { + if let Ok(delta) = serde_json::from_str::(&delta_json) { + // Stream content text to the UI + if let Some(content) = delta.get("content").and_then(|v| v.as_str()) { + if !content.is_empty() { + let mut msg = Message::assistant().with_text(content); + msg.id = Some(message_id.to_string()); + if tx.blocking_send(Ok((Some(msg), None))).is_err() { + return Ok(TokenAction::Stop); + } + } + } + // Accumulate tool call deltas + if let Some(tool_calls) = + delta.get("tool_calls").and_then(|v| v.as_array()) + { + for tc in tool_calls { + accumulated_tool_calls.push(tc.clone()); + } + } + } } } - streamed_len = stream_up_to; + Err(e) => { + tracing::warn!("Streaming parser error: {}", e); + } } let should_stop = template_result @@ -152,35 +172,29 @@ pub(super) fn generate_with_native_tools( }, )?; - let (content, tool_call_msgs) = - if let Some((xml_content, xml_calls)) = split_content_and_xml_tool_calls(&generated_text) { - let msgs = extract_xml_tool_call_messages(xml_calls, message_id); - (xml_content, msgs) - } else { - let (json_content, tool_calls_json) = split_content_and_tool_calls(&generated_text); - let msgs = tool_calls_json - .map(|tc| extract_tool_call_messages(&tc, message_id)) - .unwrap_or_default(); - (json_content, msgs) - }; - - if content.len() > streamed_len { - #[allow(clippy::string_slice)] - let remaining = &content[streamed_len..]; - if !remaining.is_empty() { - let mut msg = Message::assistant().with_text(remaining); - msg.id = Some(message_id.to_string()); - let _ = tx.blocking_send(Ok((Some(msg), None))); + // Finalize the streaming parser with is_partial=false + if let Ok(final_deltas) = stream_parser.update("", false) { + for delta_json in final_deltas { + if let Ok(delta) = serde_json::from_str::(&delta_json) { + if let Some(content) = delta.get("content").and_then(|v| v.as_str()) { + if !content.is_empty() { + let mut msg = Message::assistant().with_text(content); + msg.id = Some(message_id.to_string()); + let _ = tx.blocking_send(Ok((Some(msg), None))); + } + } + if let Some(tool_calls) = delta.get("tool_calls").and_then(|v| v.as_array()) { + for tc in tool_calls { + accumulated_tool_calls.push(tc.clone()); + } + } + } } } - if !tool_call_msgs.is_empty() { - for msg in tool_call_msgs { - let _ = tx.blocking_send(Ok((Some(msg), None))); - } - } else if content.is_empty() && !generated_text.is_empty() { - let mut msg = Message::assistant().with_text(&generated_text); - msg.id = Some(message_id.to_string()); + // Convert accumulated tool calls to messages + let tool_call_msgs = extract_oai_tool_call_messages(&accumulated_tool_calls, message_id); + for msg in tool_call_msgs { let _ = tx.blocking_send(Ok((Some(msg), None))); } @@ -195,3 +209,194 @@ pub(super) fn generate_with_native_tools( let _ = ctx.tx.blocking_send(Ok((None, Some(provider_usage)))); Ok(()) } + +/// Merge OpenAI streaming deltas by `index` into complete tool calls, then +/// convert to Goose Message objects. +/// +/// The streaming parser emits partial deltas like: +/// {"tool_calls": [{"index": 0, "id": "abc", "function": {"name": "shell"}}]} +/// {"tool_calls": [{"index": 0, "function": {"arguments": "{\"command\":"}}]} +/// {"tool_calls": [{"index": 0, "function": {"arguments": " \"ls\"}"}}]} +/// +/// These must be merged by `index` before extracting complete tool calls. +fn extract_oai_tool_call_messages(deltas: &[Value], message_id: &str) -> Vec { + let mut merged: std::collections::BTreeMap = + std::collections::BTreeMap::new(); + + for delta in deltas { + let index = delta.get("index").and_then(|v| v.as_u64()).unwrap_or(0); + let entry = merged + .entry(index) + .or_insert_with(|| (String::new(), String::new(), String::new())); + + if let Some(id) = delta.get("id").and_then(|v| v.as_str()) { + if !id.is_empty() { + entry.0 = id.to_string(); + } + } + if let Some(func) = delta.get("function") { + if let Some(name) = func.get("name").and_then(|v| v.as_str()) { + if !name.is_empty() { + entry.1 = name.to_string(); + } + } + if let Some(args) = func.get("arguments").and_then(|v| v.as_str()) { + entry.2.push_str(args); + } + } + } + + merged + .into_values() + .filter_map(|(id, name, args_str)| { + if name.is_empty() { + return None; + } + + let id = if id.is_empty() { + Uuid::new_v4().to_string() + } else { + id + }; + + let arguments: Option> = if args_str.is_empty() { + None + } else { + match serde_json::from_str(&args_str) { + Ok(args) => Some(args), + Err(_) => return None, + } + }; + + let tool_call = match arguments { + Some(args) => CallToolRequestParams::new(Cow::Owned(name)).with_arguments(args), + None => CallToolRequestParams::new(Cow::Owned(name)), + }; + + let mut msg = Message::assistant(); + msg.content + .push(MessageContent::tool_request(id, Ok(tool_call))); + msg.id = Some(message_id.to_string()); + Some(msg) + }) + .collect() +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn get_tool_call_name(msg: &Message) -> &str { + match &msg.content[0] { + MessageContent::ToolRequest(req) => { + let call = req.tool_call.as_ref().unwrap(); + &call.name + } + _ => panic!("Expected ToolRequest"), + } + } + + fn get_tool_call_args(msg: &Message) -> Option<&serde_json::Map> { + match &msg.content[0] { + MessageContent::ToolRequest(req) => { + let call = req.tool_call.as_ref().unwrap(); + call.arguments.as_ref() + } + _ => panic!("Expected ToolRequest"), + } + } + + #[test] + fn test_merge_streaming_deltas() { + // Simulates OpenAI streaming: name in first delta, arguments split across multiple + let deltas = vec![ + json!({"index": 0, "id": "call_1", "type": "function", "function": {"name": "developer__shell", "arguments": ""}}), + json!({"index": 0, "function": {"arguments": "{\"command\":"}}), + json!({"index": 0, "function": {"arguments": " \"ls\"}"}}), + ]; + let msgs = extract_oai_tool_call_messages(&deltas, "msg-1"); + assert_eq!(msgs.len(), 1); + assert_eq!(get_tool_call_name(&msgs[0]), "developer__shell"); + let args = get_tool_call_args(&msgs[0]).unwrap(); + assert_eq!(args.get("command").unwrap(), "ls"); + } + + #[test] + fn test_multiple_tool_calls_by_index() { + let deltas = vec![ + json!({"index": 0, "id": "call_1", "function": {"name": "developer__shell", "arguments": "{\"command\": \"ls\"}"}}), + json!({"index": 1, "id": "call_2", "function": {"name": "developer__shell", "arguments": "{\"command\": \"pwd\"}"}}), + ]; + let msgs = extract_oai_tool_call_messages(&deltas, "msg-1"); + assert_eq!(msgs.len(), 2); + let args0 = get_tool_call_args(&msgs[0]).unwrap(); + let args1 = get_tool_call_args(&msgs[1]).unwrap(); + assert_eq!(args0.get("command").unwrap(), "ls"); + assert_eq!(args1.get("command").unwrap(), "pwd"); + } + + #[test] + fn test_multiple_arguments_streamed() { + // Arguments with multiple keys streamed token by token + let deltas = vec![ + json!({"index": 0, "id": "call_1", "function": {"name": "developer__shell", "arguments": ""}}), + json!({"index": 0, "function": {"arguments": "{\"command\""}}), + json!({"index": 0, "function": {"arguments": ": \"ls -la\","}}), + json!({"index": 0, "function": {"arguments": " \"timeout\":"}}), + json!({"index": 0, "function": {"arguments": " 30}"}}), + ]; + let msgs = extract_oai_tool_call_messages(&deltas, "msg-1"); + assert_eq!(msgs.len(), 1); + let args = get_tool_call_args(&msgs[0]).unwrap(); + assert_eq!(args.get("command").unwrap(), "ls -la"); + assert_eq!(args.get("timeout").unwrap(), 30); + } + + #[test] + fn test_empty_name_skipped() { + let deltas = vec![json!({"index": 0, "function": {"name": "", "arguments": "{}"}})]; + let msgs = extract_oai_tool_call_messages(&deltas, "msg-1"); + assert!(msgs.is_empty()); + } + + #[test] + fn test_no_deltas() { + let msgs = extract_oai_tool_call_messages(&[], "msg-1"); + assert!(msgs.is_empty()); + } + + #[test] + fn test_tool_call_without_arguments() { + let deltas = vec![json!({"index": 0, "id": "call_1", "function": {"name": "some_tool"}})]; + let msgs = extract_oai_tool_call_messages(&deltas, "msg-1"); + assert_eq!(msgs.len(), 1); + assert_eq!(get_tool_call_name(&msgs[0]), "some_tool"); + assert!(get_tool_call_args(&msgs[0]).is_none()); + } + + #[test] + fn test_malformed_arguments_drops_tool_call() { + let deltas = vec![ + json!({"index": 0, "id": "call_1", "function": {"name": "developer__shell", "arguments": ""}}), + json!({"index": 0, "function": {"arguments": "{\"command\": \"rm -rf"}}), + ]; + let msgs = extract_oai_tool_call_messages(&deltas, "msg-1"); + assert!(msgs.is_empty()); + } + + #[test] + fn test_generates_id_when_missing() { + let deltas = + vec![json!({"index": 0, "function": {"name": "some_tool", "arguments": "{}"}})]; + let msgs = extract_oai_tool_call_messages(&deltas, "msg-1"); + assert_eq!(msgs.len(), 1); + assert_eq!(get_tool_call_name(&msgs[0]), "some_tool"); + match &msgs[0].content[0] { + MessageContent::ToolRequest(req) => { + assert!(!req.id.is_empty()); + } + _ => panic!("Expected ToolRequest"), + } + } +} diff --git a/crates/goose/src/providers/local_inference/local_model_registry.rs b/crates/goose/src/providers/local_inference/local_model_registry.rs index 3ef80c6a..d096824c 100644 --- a/crates/goose/src/providers/local_inference/local_model_registry.rs +++ b/crates/goose/src/providers/local_inference/local_model_registry.rs @@ -91,19 +91,60 @@ impl Default for ModelSettings { } } -/// Featured models — HuggingFace specs in "author/repo-GGUF:quantization" format. -pub const FEATURED_MODELS: &[&str] = &[ - "bartowski/Llama-3.2-1B-Instruct-GGUF:Q4_K_M", - "bartowski/Llama-3.2-3B-Instruct-GGUF:Q4_K_M", - "bartowski/Hermes-2-Pro-Mistral-7B-GGUF:Q4_K_M", - "bartowski/Mistral-Small-24B-Instruct-2501-GGUF:Q4_K_M", +pub struct FeaturedModel { + /// HuggingFace spec in "author/repo-GGUF:quantization" format. + pub spec: &'static str, + /// Whether this model's GGUF template supports native tool calling via llama.cpp. + pub native_tool_calling: bool, +} + +pub const FEATURED_MODELS: &[FeaturedModel] = &[ + FeaturedModel { + spec: "bartowski/Llama-3.2-1B-Instruct-GGUF:Q4_K_M", + native_tool_calling: false, + }, + FeaturedModel { + spec: "bartowski/Llama-3.2-3B-Instruct-GGUF:Q4_K_M", + native_tool_calling: false, + }, + FeaturedModel { + spec: "bartowski/Hermes-2-Pro-Mistral-7B-GGUF:Q4_K_M", + native_tool_calling: false, + }, + FeaturedModel { + spec: "bartowski/Mistral-Small-24B-Instruct-2501-GGUF:Q4_K_M", + native_tool_calling: false, + }, + FeaturedModel { + spec: "unsloth/gemma-4-E4B-it-GGUF:Q4_K_M", + native_tool_calling: true, + }, + FeaturedModel { + spec: "unsloth/gemma-4-26B-A4B-it-GGUF:Q4_K_M", + native_tool_calling: true, + }, ]; +pub fn default_settings_for_model(model_id: &str) -> ModelSettings { + use super::hf_models::parse_model_spec; + let native = FEATURED_MODELS.iter().any(|m| { + if let Ok((repo_id, quant)) = parse_model_spec(m.spec) { + model_id_from_repo(&repo_id, &quant) == model_id && m.native_tool_calling + } else { + false + } + }); + ModelSettings { + native_tool_calling: native, + ..ModelSettings::default() + } +} + /// Check if a model ID corresponds to a featured model. pub fn is_featured_model(model_id: &str) -> bool { use super::hf_models::parse_model_spec; - FEATURED_MODELS.iter().any(|spec| { - if let Ok((repo_id, quant)) = parse_model_spec(spec) { + FEATURED_MODELS.iter().any(|m| { + if let Ok((repo_id, quant)) = parse_model_spec(m.spec) { model_id_from_repo(&repo_id, &quant) == model_id } else { false diff --git a/crates/goose/src/providers/local_inference/tool_parsing.rs b/crates/goose/src/providers/local_inference/tool_parsing.rs index d2f57688..e1457663 100644 --- a/crates/goose/src/providers/local_inference/tool_parsing.rs +++ b/crates/goose/src/providers/local_inference/tool_parsing.rs @@ -1,8 +1,5 @@ -use crate::conversation::message::{Message, MessageContent}; -use rmcp::model::{CallToolRequestParams, Tool}; +use rmcp::model::Tool; use serde_json::{json, Value}; -use std::borrow::Cow; -use uuid::Uuid; pub(super) fn compact_tools_json(tools: &[Tool]) -> Option { let compact: Vec = tools @@ -20,527 +17,9 @@ pub(super) fn compact_tools_json(tools: &[Tool]) -> Option { serde_json::to_string(&compact).ok() } -/// Split generated text into (content, tool_calls_json). -/// Looks for the last top-level JSON object containing `"tool_calls"`. -/// Returns the text before it as content, and the JSON string if found. -#[allow(clippy::string_slice)] -pub(super) fn split_content_and_tool_calls(text: &str) -> (String, Option) { - let trimmed = text.trim_end(); - if !trimmed.ends_with('}') { - return (text.to_string(), None); - } - - // Scan backwards for the matching '{' of the final '}'. - // We only match on ASCII braces so `start` is always a char boundary. - let bytes = trimmed.as_bytes(); - let mut depth = 0i32; - let mut json_start = None; - for i in (0..bytes.len()).rev() { - match bytes[i] { - b'}' => depth += 1, - b'{' => { - depth -= 1; - if depth == 0 { - json_start = Some(i); - break; - } - } - _ => {} - } - } - - let Some(start) = json_start else { - return (text.to_string(), None); - }; - - let json_str = &trimmed[start..]; - let parsed: Value = match serde_json::from_str(json_str) { - Ok(v) => v, - Err(_) => return (text.to_string(), None), - }; - - if parsed - .get("tool_calls") - .and_then(|v| v.as_array()) - .is_none() - { - return (text.to_string(), None); - } - - let content = trimmed[..start].trim_end().to_string(); - (content, Some(json_str.to_string())) -} - -/// Return the byte length of text that is safe to stream. -/// Everything before the last unmatched top-level `{` is safe — the `{` could -/// be the start of a tool-call JSON block still being generated. -/// If all braces are balanced the entire text is safe. -pub(super) fn safe_stream_end(text: &str) -> usize { - // Hold back from the start of any incomplete tag. - // If we find an unmatched opening, nothing from that point should be streamed. - let xml_hold = text.find("").unwrap_or(text.len()); - - let bytes = text.as_bytes(); - let mut safe_end = bytes.len(); - let mut depth = 0i32; - for (i, &b) in bytes.iter().enumerate() { - match b { - b'{' => { - if depth == 0 { - safe_end = i; - } - depth += 1; - } - b'}' => { - depth -= 1; - if depth == 0 { - safe_end = i + 1; - } - } - _ => { - if depth == 0 { - safe_end = i + 1; - } - } - } - } - - // Also hold back a partial ``, hold them. - let tag = b""; - let tail_hold = { - let mut hold = safe_end; - let check_len = tag.len().min(bytes.len()); - for start in (safe_end.saturating_sub(check_len))..safe_end { - let tail = &bytes[start..safe_end]; - if tag.starts_with(tail) { - hold = start; - break; - } - } - hold - }; - - safe_end.min(xml_hold).min(tail_hold) -} - -/// Extract tool call messages from a JSON object containing "tool_calls". -/// Handles both the model's native format (name/arguments at top level) -/// and the OpenAI format (function.name/function.arguments). -pub(super) fn extract_tool_call_messages(tool_calls_json: &str, message_id: &str) -> Vec { - let parsed: Value = match serde_json::from_str(tool_calls_json) { - Ok(v) => v, - Err(_) => return vec![], - }; - - let Some(tool_calls) = parsed.get("tool_calls").and_then(|v| v.as_array()) else { - return vec![]; - }; - - let mut messages = Vec::new(); - for tc in tool_calls { - // Try OpenAI format first: {"function": {"name": ..., "arguments": ...}, "id": ...} - // Then model's native format: {"name": ..., "arguments": {...}, "id": ...} - let (name, arguments) = if let Some(func) = tc.get("function") { - let n = func.get("name").and_then(|v| v.as_str()).unwrap_or(""); - let args_str = func - .get("arguments") - .and_then(|v| v.as_str()) - .unwrap_or("{}"); - let args: Option> = serde_json::from_str(args_str).ok(); - (n.to_string(), args) - } else { - let n = tc.get("name").and_then(|v| v.as_str()).unwrap_or(""); - // Arguments may be an object directly (model format) or a string (OAI format) - let args = if let Some(obj) = tc.get("arguments").and_then(|v| v.as_object()) { - Some(obj.clone()) - } else if let Some(s) = tc.get("arguments").and_then(|v| v.as_str()) { - serde_json::from_str(s).ok() - } else { - None - }; - (n.to_string(), args) - }; - - if name.is_empty() { - continue; - } - - let id = tc - .get("id") - .and_then(|v| v.as_str()) - .map(|s| s.to_string()) - .unwrap_or_else(|| Uuid::new_v4().to_string()); - - let tool_call = CallToolRequestParams::new(Cow::Owned(name)) - .with_arguments(arguments.unwrap_or_default()); - - let mut msg = Message::assistant(); - msg.content - .push(MessageContent::tool_request(id, Ok(tool_call))); - msg.id = Some(message_id.to_string()); - messages.push(msg); - } - - messages -} - -/// Parse XML-style tool calls used by models like qwen3-coder. -/// Format: -/// ```text -/// -/// -/// value1 -/// value2 -/// -/// -/// ``` -/// Returns (content_before_tool_calls, vec_of_tool_calls) or None if no XML tool calls found. -#[allow(clippy::type_complexity)] -pub(super) fn split_content_and_xml_tool_calls( - text: &str, -) -> Option<(String, Vec<(String, serde_json::Map)>)> { - let (content, first_block_and_rest) = text.split_once("")?; - let content = content.trim_end().to_string(); - let mut tool_calls = Vec::new(); - - // Process the first block, then keep splitting on subsequent tags - let mut remaining = first_block_and_rest; - loop { - // Split off the block up to (or take the rest if unclosed) - let (block, after_close) = remaining - .split_once("") - .unwrap_or((remaining, "")); - - if let Some(tool_call) = parse_single_xml_tool_call(block) { - tool_calls.push(tool_call); - } - - // Try to find the next in what remains - match after_close.split_once("") { - Some((_between, next_remaining)) => remaining = next_remaining, - None => break, - } - } - - if tool_calls.is_empty() { - None - } else { - Some((content, tool_calls)) - } -} - -fn parse_single_xml_tool_call(block: &str) -> Option<(String, serde_json::Map)> { - // Try V... format first - if let Some(result) = parse_xml_function_format(block) { - return Some(result); - } - // Try GLM-style: TOOL_NAMEKV... - parse_xml_arg_key_value_format(block) -} - -fn parse_xml_function_format(block: &str) -> Option<(String, serde_json::Map)> { - let (_, after_func_eq) = block.split_once("')?; - let func_name = func_name.trim().to_string(); - - let mut args = serde_json::Map::new(); - let mut rest = func_body; - - while let Some((_, after_param_eq)) = rest.split_once("') else { - break; - }; - let param_name = param_name.trim().to_string(); - - let (value, after_value) = after_name_close - .split_once("") - .unwrap_or((after_name_close, "")); - let value = value.trim(); - - let json_value = - serde_json::from_str(value).unwrap_or_else(|_| Value::String(value.to_string())); - args.insert(param_name, json_value); - - rest = after_value; - } - - Some((func_name, args)) -} - -/// Parse GLM-style tool calls: `NAMEKV...` -/// Also handles zero-argument calls like just `NAME`. -fn parse_xml_arg_key_value_format(block: &str) -> Option<(String, serde_json::Map)> { - let func_name_end = block.find("").unwrap_or(block.len()); - // Safe: find returns a byte offset at the start of an ASCII '<' character, - // and block.len() is always a valid boundary. - #[allow(clippy::string_slice)] - let func_name = block[..func_name_end].trim().to_string(); - if func_name.is_empty() { - return None; - } - - let mut args = serde_json::Map::new(); - #[allow(clippy::string_slice)] - let mut rest = &block[func_name_end..]; - - while let Some((_, after_key_open)) = rest.split_once("") { - let Some((key, after_key_close)) = after_key_open.split_once("") else { - break; - }; - let key = key.trim().to_string(); - - let Some((_, after_val_open)) = after_key_close.split_once("") else { - break; - }; - let (value, after_val_close) = after_val_open - .split_once("") - .unwrap_or((after_val_open, "")); - let value = value.trim(); - - let json_value = - serde_json::from_str(value).unwrap_or_else(|_| Value::String(value.to_string())); - args.insert(key, json_value); - - rest = after_val_close; - } - - Some((func_name, args)) -} - -pub(super) fn extract_xml_tool_call_messages( - tool_calls: Vec<(String, serde_json::Map)>, - message_id: &str, -) -> Vec { - tool_calls - .into_iter() - .map(|(name, args)| { - let tool_call = if args.is_empty() { - CallToolRequestParams::new(Cow::Owned(name)) - } else { - CallToolRequestParams::new(Cow::Owned(name)).with_arguments(args) - }; - let mut msg = Message::assistant(); - msg.content.push(MessageContent::tool_request( - Uuid::new_v4().to_string(), - Ok(tool_call), - )); - msg.id = Some(message_id.to_string()); - msg - }) - .collect() -} - #[cfg(test)] mod tests { use super::*; - use serde_json::json; - - const SHELL_TOOL: &str = "developer__shell"; - - #[test] - fn test_parse_xml_tool_call_single() { - let text = "I'll search for that.\n\n\n\nlocal.*inference\n\n"; - let result = split_content_and_xml_tool_calls(text); - assert!(result.is_some()); - let (content, calls) = result.unwrap(); - assert_eq!(content, "I'll search for that."); - assert_eq!(calls.len(), 1); - assert_eq!(calls[0].0, "search__files"); - assert_eq!(calls[0].1.get("pattern").unwrap(), "local.*inference"); - } - - #[test] - fn test_parse_xml_tool_call_multiple_params() { - let text = "\n\nls -la\n30\n\n"; - let result = split_content_and_xml_tool_calls(text); - assert!(result.is_some()); - let (content, calls) = result.unwrap(); - assert!(content.is_empty()); - assert_eq!(calls.len(), 1); - assert_eq!(calls[0].0, SHELL_TOOL); - assert_eq!(calls[0].1.get("command").unwrap(), "ls -la"); - // 30 should be parsed as a number - assert_eq!(calls[0].1.get("timeout").unwrap(), &json!(30)); - } - - #[test] - fn test_parse_xml_tool_call_no_tool_call() { - let text = "Just some regular text with no tool calls."; - assert!(split_content_and_xml_tool_calls(text).is_none()); - } - - #[test] - fn test_parse_xml_tool_call_multiple_calls() { - let text = "Doing two things:\n\n\n1\n\n\n\n\nhello\n\n"; - let result = split_content_and_xml_tool_calls(text); - assert!(result.is_some()); - let (content, calls) = result.unwrap(); - assert_eq!(content, "Doing two things:"); - assert_eq!(calls.len(), 2); - assert_eq!(calls[0].0, "foo__bar"); - assert_eq!(calls[1].0, "baz__qux"); - } - - #[test] - fn test_parse_xml_tool_call_multiline_value() { - let text = "\n\ntest.py\ndef hello():\n print(\"world\")\n\n"; - let result = split_content_and_xml_tool_calls(text); - assert!(result.is_some()); - let (_content, calls) = result.unwrap(); - assert_eq!(calls[0].0, "developer__write_file"); - assert_eq!( - calls[0].1.get("content").unwrap(), - "def hello():\n print(\"world\")" - ); - } - - #[test] - fn test_safe_stream_end_holds_back_tool_call_tag() { - let text = "Some text before \n"; - let safe = safe_stream_end(text); - assert!(safe <= text.find("").unwrap()); - } - - #[test] - fn test_safe_stream_end_holds_back_partial_tag() { - let text = "Some text "; - let result = split_content_and_xml_tool_calls(text); - assert!(result.is_some()); - let (content, calls) = result.unwrap(); - assert_eq!(content, "Let me check."); - assert_eq!(calls.len(), 1); - assert_eq!(calls[0].0, "execute_typescript"); - assert_eq!( - calls[0].1.get("code").unwrap(), - "async function run() { return 1; }" - ); - // tool_graph should be parsed as JSON array - assert!(calls[0].1.get("tool_graph").unwrap().is_array()); - } - - #[test] - fn test_extract_xml_tool_call_messages() { - let calls = vec![( - SHELL_TOOL.to_string(), - serde_json::Map::from_iter(vec![("command".to_string(), json!("ls"))]), - )]; - let msgs = extract_xml_tool_call_messages(calls, "test-id"); - assert_eq!(msgs.len(), 1); - assert_eq!(msgs[0].id, Some("test-id".to_string())); - match &msgs[0].content[0] { - MessageContent::ToolRequest(req) => { - let call = req.tool_call.as_ref().unwrap(); - assert_eq!(&*call.name, SHELL_TOOL); - assert_eq!( - call.arguments.as_ref().unwrap().get("command").unwrap(), - "ls" - ); - } - _ => panic!("Expected ToolRequest"), - } - } - - #[test] - fn test_split_content_and_tool_calls_with_tool() { - let text = "Here is the result.\n{\"tool_calls\": [{\"function\": {\"name\": \"shell\", \"arguments\": \"{}\"}, \"id\": \"abc\"}]}"; - let (content, tc) = split_content_and_tool_calls(text); - assert_eq!(content, "Here is the result."); - assert!(tc.is_some()); - let parsed: serde_json::Value = serde_json::from_str(&tc.unwrap()).unwrap(); - assert_eq!(parsed["tool_calls"].as_array().unwrap().len(), 1); - } - - #[test] - fn test_split_content_and_tool_calls_no_tool() { - let text = "Just regular text, no JSON."; - let (content, tc) = split_content_and_tool_calls(text); - assert_eq!(content, text); - assert!(tc.is_none()); - } - - #[test] - fn test_split_content_and_tool_calls_json_without_tool_calls_key() { - let text = "{\"key\": \"value\"}"; - let (content, tc) = split_content_and_tool_calls(text); - assert_eq!(content, text); - assert!(tc.is_none()); - } - - #[test] - fn test_extract_tool_call_messages_openai_format() { - let json_str = r#"{"tool_calls": [{"function": {"name": "developer__shell", "arguments": "{\"command\": \"ls\"}"}, "id": "call-1"}]}"#; - let msgs = extract_tool_call_messages(json_str, "msg-1"); - assert_eq!(msgs.len(), 1); - match &msgs[0].content[0] { - MessageContent::ToolRequest(req) => { - let call = req.tool_call.as_ref().unwrap(); - assert_eq!(&*call.name, SHELL_TOOL); - assert_eq!( - call.arguments.as_ref().unwrap().get("command").unwrap(), - "ls" - ); - } - _ => panic!("Expected ToolRequest"), - } - } - - #[test] - fn test_extract_tool_call_messages_native_format() { - let json_str = r#"{"tool_calls": [{"name": "developer__shell", "arguments": {"command": "ls"}, "id": "call-2"}]}"#; - let msgs = extract_tool_call_messages(json_str, "msg-2"); - assert_eq!(msgs.len(), 1); - match &msgs[0].content[0] { - MessageContent::ToolRequest(req) => { - let call = req.tool_call.as_ref().unwrap(); - assert_eq!(&*call.name, SHELL_TOOL); - } - _ => panic!("Expected ToolRequest"), - } - } - - #[test] - fn test_extract_tool_call_messages_invalid_json() { - assert!(extract_tool_call_messages("not json", "msg-3").is_empty()); - } - - #[test] - fn test_extract_tool_call_messages_empty_name_skipped() { - let json_str = r#"{"tool_calls": [{"name": "", "arguments": {}, "id": "x"}]}"#; - assert!(extract_tool_call_messages(json_str, "msg-4").is_empty()); - } #[test] fn test_compact_tools_json_produces_minimal_output() { @@ -570,27 +49,4 @@ mod tests { let parsed: Vec = serde_json::from_str(&result.unwrap()).unwrap(); assert!(parsed.is_empty()); } - - #[test] - fn test_safe_stream_end_balanced_braces() { - let text = "Result: {\"key\": \"value\"} done"; - assert_eq!(safe_stream_end(text), text.len()); - } - - #[test] - fn test_safe_stream_end_unbalanced_open_brace() { - let text = "Some text {\"tool_calls\": ["; - assert_eq!(safe_stream_end(text), "Some text ".len()); - } - - #[test] - fn test_safe_stream_end_empty() { - assert_eq!(safe_stream_end(""), 0); - } - - #[test] - fn test_safe_stream_end_no_braces() { - let text = "plain text here"; - assert_eq!(safe_stream_end(text), text.len()); - } } diff --git a/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx b/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx index f5791823..959eba07 100644 --- a/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx +++ b/ui/desktop/src/components/settings/models/bottom_bar/ModelsBottomBar.tsx @@ -162,7 +162,7 @@ export default function ModelsBottomBar({ {isLocalModelSettingsOpen && currentModel && (
-
+

{intl.formatMessage(i18n.localModelSettingsTitle, { modelName: getModelDisplayName(currentModel) })}