fix: local inference provider loses thinking and fragments tool calls across turns (#8415)

Signed-off-by: jh-block <jhugo@block.xyz>
This commit is contained in:
jh-block
2026-04-09 11:27:31 +02:00
committed by GitHub
parent dbd64a7297
commit e8cbfc55da
5 changed files with 87 additions and 50 deletions
@@ -485,6 +485,14 @@ impl Provider for LocalInferenceProvider {
}
}
// Allow request_params to override thinking
let mut model_settings = model_settings;
if let Some(false) =
model_config.get_config_param::<bool>("enable_thinking", "GOOSE_LOCAL_ENABLE_THINKING")
{
model_settings.enable_thinking = false;
}
// 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;
@@ -33,12 +33,16 @@ pub(super) fn generate_with_native_tools(
tool_choice: None,
json_schema: None,
grammar: None,
reasoning_format: Some("auto"),
reasoning_format: if ctx.settings.enable_thinking {
Some("auto")
} else {
None
},
chat_template_kwargs: None,
add_generation_prompt: true,
use_jinja: true,
parallel_tool_calls: false,
enable_thinking: true,
enable_thinking: ctx.settings.enable_thinking,
add_bos: false,
add_eos: false,
parse_tool_calls: true,
@@ -119,6 +123,10 @@ pub(super) fn generate_with_native_tools(
// Accumulate tool calls across streaming deltas
let mut accumulated_tool_calls: Vec<Value> = Vec::new();
// Accumulate thinking/reasoning across the entire generation so we can
// attach it to the final tool-call message (mirroring what the OpenAI
// streaming path does). Streaming chunks are still sent for UI display.
let mut accumulated_thinking = String::new();
let output_token_count = generation_loop(
&ctx.loaded.model,
@@ -139,6 +147,7 @@ pub(super) fn generate_with_native_tools(
delta.get("reasoning_content").and_then(|v| v.as_str())
{
if !reasoning.is_empty() {
accumulated_thinking.push_str(reasoning);
let mut msg = Message::assistant().with_thinking(reasoning, "");
msg.id = Some(message_id.to_string());
if tx.blocking_send(Ok((Some(msg), None))).is_err() {
@@ -195,6 +204,7 @@ pub(super) fn generate_with_native_tools(
if let Ok(delta) = serde_json::from_str::<Value>(&delta_json) {
if let Some(reasoning) = delta.get("reasoning_content").and_then(|v| v.as_str()) {
if !reasoning.is_empty() {
accumulated_thinking.push_str(reasoning);
let mut msg = Message::assistant().with_thinking(reasoning, "");
msg.id = Some(message_id.to_string());
let _ = tx.blocking_send(Ok((Some(msg), None)));
@@ -216,9 +226,25 @@ pub(super) fn generate_with_native_tools(
}
}
// 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 {
// Build a single message combining thinking + all tool calls, mirroring
// the structure produced by the OpenAI streaming path. The agent relies
// on this combined message to:
// 1. Extract thinking and attach it to per-tool-request messages
// 2. Enable merge_split_tool_call_messages to reconstruct the standard
// OpenAI format (one assistant msg with N tool_calls, then N tool results)
let tool_call_contents = extract_oai_tool_call_contents(&accumulated_tool_calls);
if !tool_call_contents.is_empty() {
let mut contents: Vec<MessageContent> = Vec::new();
if !accumulated_thinking.is_empty() {
contents.push(MessageContent::thinking(&accumulated_thinking, ""));
}
contents.extend(tool_call_contents);
let mut msg = Message::new(
rmcp::model::Role::Assistant,
chrono::Utc::now().timestamp(),
contents,
);
msg.id = Some(message_id.to_string());
let _ = tx.blocking_send(Ok((Some(msg), None)));
}
@@ -234,16 +260,12 @@ pub(super) fn generate_with_native_tools(
Ok(())
}
/// Merge OpenAI streaming deltas by `index` into complete tool calls, then
/// convert to Goose Message objects.
/// Merge OpenAI streaming deltas by `index` into `MessageContent` items.
///
/// 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<Message> {
/// Returns one `ToolRequest` content per distinct tool call index. The caller
/// is responsible for combining these into a single `Message` (together with
/// any accumulated thinking content).
fn extract_oai_tool_call_contents(deltas: &[Value]) -> Vec<MessageContent> {
let mut merged: std::collections::BTreeMap<u64, (String, String, String)> =
std::collections::BTreeMap::new();
@@ -297,11 +319,7 @@ fn extract_oai_tool_call_messages(deltas: &[Value], message_id: &str) -> Vec<Mes
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)
Some(MessageContent::tool_request(id, Ok(tool_call)))
})
.collect()
}
@@ -311,8 +329,8 @@ mod tests {
use super::*;
use serde_json::json;
fn get_tool_call_name(msg: &Message) -> &str {
match &msg.content[0] {
fn get_content_tool_call_name(content: &MessageContent) -> &str {
match content {
MessageContent::ToolRequest(req) => {
let call = req.tool_call.as_ref().unwrap();
&call.name
@@ -321,8 +339,10 @@ mod tests {
}
}
fn get_tool_call_args(msg: &Message) -> Option<&serde_json::Map<String, Value>> {
match &msg.content[0] {
fn get_content_tool_call_args(
content: &MessageContent,
) -> Option<&serde_json::Map<String, Value>> {
match content {
MessageContent::ToolRequest(req) => {
let call = req.tool_call.as_ref().unwrap();
call.arguments.as_ref()
@@ -333,16 +353,15 @@ mod tests {
#[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();
let contents = extract_oai_tool_call_contents(&deltas);
assert_eq!(contents.len(), 1);
assert_eq!(get_content_tool_call_name(&contents[0]), "developer__shell");
let args = get_content_tool_call_args(&contents[0]).unwrap();
assert_eq!(args.get("command").unwrap(), "ls");
}
@@ -352,17 +371,16 @@ mod tests {
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();
let contents = extract_oai_tool_call_contents(&deltas);
assert_eq!(contents.len(), 2);
let args0 = get_content_tool_call_args(&contents[0]).unwrap();
let args1 = get_content_tool_call_args(&contents[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\""}}),
@@ -370,9 +388,9 @@ mod tests {
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();
let contents = extract_oai_tool_call_contents(&deltas);
assert_eq!(contents.len(), 1);
let args = get_content_tool_call_args(&contents[0]).unwrap();
assert_eq!(args.get("command").unwrap(), "ls -la");
assert_eq!(args.get("timeout").unwrap(), 30);
}
@@ -380,23 +398,23 @@ mod tests {
#[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());
let contents = extract_oai_tool_call_contents(&deltas);
assert!(contents.is_empty());
}
#[test]
fn test_no_deltas() {
let msgs = extract_oai_tool_call_messages(&[], "msg-1");
assert!(msgs.is_empty());
let contents = extract_oai_tool_call_contents(&[]);
assert!(contents.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());
let contents = extract_oai_tool_call_contents(&deltas);
assert_eq!(contents.len(), 1);
assert_eq!(get_content_tool_call_name(&contents[0]), "some_tool");
assert!(get_content_tool_call_args(&contents[0]).is_none());
}
#[test]
@@ -405,18 +423,18 @@ mod tests {
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());
let contents = extract_oai_tool_call_contents(&deltas);
assert!(contents.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] {
let contents = extract_oai_tool_call_contents(&deltas);
assert_eq!(contents.len(), 1);
assert_eq!(get_content_tool_call_name(&contents[0]), "some_tool");
match &contents[0] {
MessageContent::ToolRequest(req) => {
assert!(!req.id.is_empty());
}
@@ -60,6 +60,12 @@ pub struct ModelSettings {
pub native_tool_calling: bool,
#[serde(default)]
pub use_jinja: bool,
#[serde(default = "default_true")]
pub enable_thinking: bool,
}
fn default_true() -> bool {
true
}
fn default_repeat_penalty() -> f32 {
@@ -87,6 +93,7 @@ impl Default for ModelSettings {
n_threads: None,
native_tool_calling: false,
use_jinja: false,
enable_thinking: true,
}
}
}
+3
View File
@@ -6486,6 +6486,9 @@
"nullable": true,
"minimum": 0
},
"enable_thinking": {
"type": "boolean"
},
"flash_attention": {
"type": "boolean",
"nullable": true
+1
View File
@@ -818,6 +818,7 @@ export type ModelInfoResponse = {
export type ModelSettings = {
context_size?: number | null;
enable_thinking?: boolean;
flash_attention?: boolean | null;
frequency_penalty?: number;
max_output_tokens?: number | null;