From 5c17964784f39755609c2acf692d609155241bce Mon Sep 17 00:00:00 2001 From: Brandon Kvarda Date: Thu, 12 Feb 2026 03:48:05 -0800 Subject: [PATCH] feat: support extra field in chatcompletion tool_calls for gemini openai compat (#6184) Signed-off-by: Brandon Kvarda Co-authored-by: Brandon Kvarda Co-authored-by: Claude Opus 4.5 --- .../goose/src/providers/formats/databricks.rs | 92 ++++++++- crates/goose/src/providers/formats/openai.rs | 194 ++++++++++++++++-- 2 files changed, 268 insertions(+), 18 deletions(-) diff --git a/crates/goose/src/providers/formats/databricks.rs b/crates/goose/src/providers/formats/databricks.rs index a55d26f72f..2243794b70 100644 --- a/crates/goose/src/providers/formats/databricks.rs +++ b/crates/goose/src/providers/formats/databricks.rs @@ -161,11 +161,23 @@ fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec { content_array @@ -1382,6 +1394,39 @@ mod tests { Ok(()) } + #[test] + fn test_format_messages_with_thought_signature_metadata() -> anyhow::Result<()> { + let mut metadata = serde_json::Map::new(); + metadata.insert( + "thoughtSignature".to_string(), + json!("sig_abc123_test_signature"), + ); + + let message = Message::assistant().with_tool_request_with_metadata( + "tool1", + Ok(CallToolRequestParams { + meta: None, + task: None, + name: "test_tool".into(), + arguments: Some(object!({"param": "value"})), + }), + Some(&metadata), + None, + ); + + let spec = format_messages(&[message], &ImageFormat::OpenAi); + let as_value = serde_json::to_value(spec)?; + let spec_array = as_value.as_array().unwrap(); + + assert_eq!(spec_array.len(), 1); + let tool_call = &spec_array[0]["tool_calls"][0]; + assert_eq!(tool_call["id"], "tool1"); + assert_eq!(tool_call["function"]["name"], "test_tool"); + assert_eq!(tool_call["thoughtSignature"], "sig_abc123_test_signature"); + + Ok(()) + } + #[test] fn test_create_request_claude_has_cache_control() -> anyhow::Result<()> { let model_config = ModelConfig { @@ -1477,4 +1522,45 @@ mod tests { Ok(()) } + + #[test] + fn test_format_messages_with_multiple_metadata_fields() -> anyhow::Result<()> { + let mut metadata = serde_json::Map::new(); + metadata.insert("thoughtSignature".to_string(), json!("sig_top_level")); + metadata.insert( + "extra_content".to_string(), + json!({ + "google": { + "thought_signature": "sig_nested" + } + }), + ); + metadata.insert("custom_field".to_string(), json!("custom_value")); + + let message = Message::assistant().with_tool_request_with_metadata( + "tool1", + Ok(CallToolRequestParams { + meta: None, + task: None, + name: "test_tool".into(), + arguments: None, + }), + Some(&metadata), + None, + ); + + let spec = format_messages(&[message], &ImageFormat::OpenAi); + let as_value = serde_json::to_value(spec)?; + let spec_array = as_value.as_array().unwrap(); + + let tool_call = &spec_array[0]["tool_calls"][0]; + assert_eq!(tool_call["thoughtSignature"], "sig_top_level"); + assert_eq!( + tool_call["extra_content"]["google"]["thought_signature"], + "sig_nested" + ); + assert_eq!(tool_call["custom_field"], "custom_value"); + + Ok(()) + } } diff --git a/crates/goose/src/providers/formats/openai.rs b/crates/goose/src/providers/formats/openai.rs index 618b8d8f40..3103527305 100644 --- a/crates/goose/src/providers/formats/openai.rs +++ b/crates/goose/src/providers/formats/openai.rs @@ -16,8 +16,19 @@ use rmcp::model::{ use serde::{Deserialize, Serialize}; use serde_json::{json, Value}; use std::borrow::Cow; +use std::collections::HashMap; use std::ops::Deref; +type ToolCallData = HashMap< + i32, + ( + String, + String, + String, + Option>, + ), +>; + #[derive(Serialize, Deserialize, Debug, Default)] struct DeltaToolCallFunction { name: Option, @@ -31,6 +42,8 @@ struct DeltaToolCall { function: DeltaToolCallFunction, index: Option, r#type: Option, + #[serde(flatten)] + extra: Option>, } #[derive(Serialize, Deserialize, Debug)] @@ -115,14 +128,22 @@ pub fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec< .entry("tool_calls") .or_insert(json!([])); - tool_calls.as_array_mut().unwrap().push(json!({ + let mut tool_call_json = json!({ "id": request.id, "type": "function", "function": { "name": sanitized_name, "arguments": arguments_str, } - })); + }); + + if let Some(metadata) = &request.metadata { + for (key, value) in metadata { + tool_call_json[key] = value.clone(); + } + } + + tool_calls.as_array_mut().unwrap().push(tool_call_json); } Err(e) => { output.push(json!({ @@ -337,6 +358,17 @@ pub fn response_to_message(response: &Value) -> anyhow::Result { arguments_str }; + let standard_fields = ["id", "function", "type", "index"]; + let metadata: Option> = tool_call + .as_object() + .map(|obj| { + obj.iter() + .filter(|(k, _)| !standard_fields.contains(&k.as_str())) + .map(|(k, v)| (k.clone(), v.clone())) + .collect() + }) + .filter(|m: &serde_json::Map| !m.is_empty()); + if !is_valid_function_name(&function_name) { let error = ErrorData { code: ErrorCode::INVALID_REQUEST, @@ -346,11 +378,15 @@ pub fn response_to_message(response: &Value) -> anyhow::Result { )), data: None, }; - content.push(MessageContent::tool_request(id, Err(error))); + content.push(MessageContent::tool_request_with_metadata( + id, + Err(error), + metadata.as_ref(), + )); } else { match safely_parse_json(&arguments_str) { Ok(params) => { - content.push(MessageContent::tool_request( + content.push(MessageContent::tool_request_with_metadata( id, Ok(CallToolRequestParams { meta: None, @@ -358,6 +394,7 @@ pub fn response_to_message(response: &Value) -> anyhow::Result { name: function_name.into(), arguments: Some(object(params)), }), + metadata.as_ref(), )); } Err(e) => { @@ -369,7 +406,11 @@ pub fn response_to_message(response: &Value) -> anyhow::Result { )), data: None, }; - content.push(MessageContent::tool_request(id, Err(error))); + content.push(MessageContent::tool_request_with_metadata( + id, + Err(error), + metadata.as_ref(), + )); } } } @@ -507,12 +548,12 @@ where if chunk.choices.is_empty() { yield (None, usage) } else if chunk.choices[0].delta.tool_calls.as_ref().is_some_and(|tc| !tc.is_empty()) { - let mut tool_call_data: std::collections::HashMap = std::collections::HashMap::new(); + let mut tool_call_data: ToolCallData = HashMap::new(); if let Some(tool_calls) = &chunk.choices[0].delta.tool_calls { for tool_call in tool_calls { if let (Some(index), Some(id), Some(name)) = (tool_call.index, &tool_call.id, &tool_call.function.name) { - tool_call_data.insert(index, (id.clone(), name.clone(), tool_call.function.arguments.clone())); + tool_call_data.insert(index, (id.clone(), name.clone(), tool_call.function.arguments.clone(), tool_call.extra.clone())); } } } @@ -543,10 +584,17 @@ where if let Some(delta_tool_calls) = &tool_chunk.choices[0].delta.tool_calls { for delta_call in delta_tool_calls { if let Some(index) = delta_call.index { - if let Some((_, _, ref mut args)) = tool_call_data.get_mut(&index) { + if let Some((_, _, ref mut args, ref mut extra)) = tool_call_data.get_mut(&index) { args.push_str(&delta_call.function.arguments); + if extra.is_none() && delta_call.extra.is_some() { + *extra = delta_call.extra.clone(); + } else if let (Some(existing), Some(new_extra)) = (extra.as_mut(), &delta_call.extra) { + for (key, value) in new_extra { + existing.entry(key.clone()).or_insert(value.clone()); + } + } } else if let (Some(id), Some(name)) = (&delta_call.id, &delta_call.function.name) { - tool_call_data.insert(index, (id.clone(), name.clone(), delta_call.function.arguments.clone())); + tool_call_data.insert(index, (id.clone(), name.clone(), delta_call.function.arguments.clone(), delta_call.extra.clone())); } } } @@ -564,7 +612,7 @@ where } } - let metadata: Option = if !accumulated_reasoning.is_empty() { + let _metadata: Option = if !accumulated_reasoning.is_empty() { let mut map = ProviderMetadata::new(); map.insert("reasoning_details".to_string(), json!(accumulated_reasoning)); Some(map) @@ -577,23 +625,26 @@ where sorted_indices.sort(); for index in sorted_indices { - if let Some((id, function_name, arguments)) = tool_call_data.get(&index) { + if let Some((id, function_name, arguments, extra_fields)) = tool_call_data.get(&index) { let parsed = if arguments.is_empty() { Ok(json!({})) } else { serde_json::from_str::(arguments) }; + let metadata = extra_fields.as_ref().filter(|m| !m.is_empty()); + let content = match parsed { Ok(params) => { MessageContent::tool_request_with_metadata( id.clone(), Ok(CallToolRequestParams { - meta: None, task: None, + meta: None, + task: None, name: function_name.clone().into(), - arguments: Some(object(params)) + arguments: Some(object(params)), }), - metadata.as_ref(), + metadata, ) }, Err(e) => { @@ -605,7 +656,7 @@ where )), data: None, }; - MessageContent::tool_request_with_metadata(id.clone(), Err(error), metadata.as_ref()) + MessageContent::tool_request_with_metadata(id.clone(), Err(error), metadata) } }; contents.push(content); @@ -1670,4 +1721,117 @@ data: [DONE] Ok(()) } + + #[test] + fn test_response_to_message_with_nested_extra_content() -> anyhow::Result<()> { + let response = json!({ + "choices": [{ + "message": { + "role": "assistant", + "tool_calls": [{ + "id": "call_456", + "type": "function", + "function": { + "name": "test_tool", + "arguments": "{}" + }, + "extra_content": { + "google": { + "thought_signature": "nested_sig_xyz789" + } + } + }] + } + }] + }); + + let message = response_to_message(&response)?; + assert_eq!(message.content.len(), 1); + + if let MessageContent::ToolRequest(request) = &message.content[0] { + assert!(request.tool_call.is_ok()); + assert!(request.metadata.is_some()); + let metadata = request.metadata.as_ref().unwrap(); + let extra_content = metadata.get("extra_content").unwrap(); + assert_eq!( + extra_content["google"]["thought_signature"], + "nested_sig_xyz789" + ); + } else { + panic!("Expected ToolRequest content"); + } + + Ok(()) + } + + #[test] + fn test_response_to_message_with_multiple_extra_fields() -> anyhow::Result<()> { + let response = json!({ + "choices": [{ + "message": { + "role": "assistant", + "tool_calls": [{ + "id": "call_789", + "type": "function", + "function": { + "name": "test_tool", + "arguments": "{}" + }, + "thoughtSignature": "sig_top_level", + "extra_content": { + "google": { + "thought_signature": "sig_nested" + } + }, + "custom_field": "custom_value" + }] + } + }] + }); + + let message = response_to_message(&response)?; + + if let MessageContent::ToolRequest(request) = &message.content[0] { + let metadata = request.metadata.as_ref().unwrap(); + assert_eq!(metadata.get("thoughtSignature").unwrap(), "sig_top_level"); + assert_eq!( + metadata.get("extra_content").unwrap()["google"]["thought_signature"], + "sig_nested" + ); + assert_eq!(metadata.get("custom_field").unwrap(), "custom_value"); + } else { + panic!("Expected ToolRequest content"); + } + + Ok(()) + } + + #[tokio::test] + async fn test_streaming_response_with_nested_extra_content() -> anyhow::Result<()> { + let response_lines = r#"data: {"model":"test-model","choices":[{"delta":{"role":"assistant","tool_calls":[{"extra_content":{"google":{"thought_signature":"nested_stream_sig"}},"id":"call_nested","function":{"name":"test_tool","arguments":"{}"},"type":"function","index":0}]},"index":0,"finish_reason":"tool_calls"}],"usage":{"prompt_tokens":100,"completion_tokens":10,"total_tokens":110},"object":"chat.completion.chunk","id":"test-id","created":1234567890} +data: [DONE]"#; + + let response_stream = + tokio_stream::iter(response_lines.lines().map(|line| Ok(line.to_string()))); + let messages = response_to_streaming_message(response_stream); + pin!(messages); + + while let Some(Ok((message, _usage))) = messages.next().await { + if let Some(msg) = message { + if let MessageContent::ToolRequest(request) = &msg.content[0] { + assert!(request.tool_call.is_ok()); + assert!(request.metadata.is_some()); + let metadata = request.metadata.as_ref().unwrap(); + let extra_content = metadata.get("extra_content").unwrap(); + assert_eq!( + extra_content["google"]["thought_signature"], + "nested_stream_sig" + ); + return Ok(()); + } + } + } + + panic!("Expected tool call message with nested extra_content metadata"); + } }