feat: support extra field in chatcompletion tool_calls for gemini openai compat (#6184)
Signed-off-by: Brandon Kvarda <brandon.kvarda@databricks.com> Co-authored-by: Brandon Kvarda <brandon.kvarda@databricks.com> Co-authored-by: Claude Opus 4.5 <noreply@anthropic.com>
This commit is contained in:
@@ -161,11 +161,23 @@ fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<Data
|
|||||||
})
|
})
|
||||||
.unwrap_or_else(|| "{}".to_string());
|
.unwrap_or_else(|| "{}".to_string());
|
||||||
|
|
||||||
converted.tool_calls.get_or_insert_default().push(json!({
|
let tool_calls = converted.tool_calls.get_or_insert_default();
|
||||||
|
let mut tool_call_json = json!({
|
||||||
"id": request.id,
|
"id": request.id,
|
||||||
"type": "function",
|
"type": "function",
|
||||||
"function": {"name": sanitized_name, "arguments": arguments_str}
|
"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.push(tool_call_json);
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
content_array
|
content_array
|
||||||
@@ -1382,6 +1394,39 @@ mod tests {
|
|||||||
Ok(())
|
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]
|
#[test]
|
||||||
fn test_create_request_claude_has_cache_control() -> anyhow::Result<()> {
|
fn test_create_request_claude_has_cache_control() -> anyhow::Result<()> {
|
||||||
let model_config = ModelConfig {
|
let model_config = ModelConfig {
|
||||||
@@ -1477,4 +1522,45 @@ mod tests {
|
|||||||
|
|
||||||
Ok(())
|
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(())
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -16,8 +16,19 @@ use rmcp::model::{
|
|||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use serde_json::{json, Value};
|
use serde_json::{json, Value};
|
||||||
use std::borrow::Cow;
|
use std::borrow::Cow;
|
||||||
|
use std::collections::HashMap;
|
||||||
use std::ops::Deref;
|
use std::ops::Deref;
|
||||||
|
|
||||||
|
type ToolCallData = HashMap<
|
||||||
|
i32,
|
||||||
|
(
|
||||||
|
String,
|
||||||
|
String,
|
||||||
|
String,
|
||||||
|
Option<serde_json::Map<String, Value>>,
|
||||||
|
),
|
||||||
|
>;
|
||||||
|
|
||||||
#[derive(Serialize, Deserialize, Debug, Default)]
|
#[derive(Serialize, Deserialize, Debug, Default)]
|
||||||
struct DeltaToolCallFunction {
|
struct DeltaToolCallFunction {
|
||||||
name: Option<String>,
|
name: Option<String>,
|
||||||
@@ -31,6 +42,8 @@ struct DeltaToolCall {
|
|||||||
function: DeltaToolCallFunction,
|
function: DeltaToolCallFunction,
|
||||||
index: Option<i32>,
|
index: Option<i32>,
|
||||||
r#type: Option<String>,
|
r#type: Option<String>,
|
||||||
|
#[serde(flatten)]
|
||||||
|
extra: Option<serde_json::Map<String, Value>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
#[derive(Serialize, Deserialize, Debug)]
|
#[derive(Serialize, Deserialize, Debug)]
|
||||||
@@ -115,14 +128,22 @@ pub fn format_messages(messages: &[Message], image_format: &ImageFormat) -> Vec<
|
|||||||
.entry("tool_calls")
|
.entry("tool_calls")
|
||||||
.or_insert(json!([]));
|
.or_insert(json!([]));
|
||||||
|
|
||||||
tool_calls.as_array_mut().unwrap().push(json!({
|
let mut tool_call_json = json!({
|
||||||
"id": request.id,
|
"id": request.id,
|
||||||
"type": "function",
|
"type": "function",
|
||||||
"function": {
|
"function": {
|
||||||
"name": sanitized_name,
|
"name": sanitized_name,
|
||||||
"arguments": arguments_str,
|
"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) => {
|
Err(e) => {
|
||||||
output.push(json!({
|
output.push(json!({
|
||||||
@@ -337,6 +358,17 @@ pub fn response_to_message(response: &Value) -> anyhow::Result<Message> {
|
|||||||
arguments_str
|
arguments_str
|
||||||
};
|
};
|
||||||
|
|
||||||
|
let standard_fields = ["id", "function", "type", "index"];
|
||||||
|
let metadata: Option<serde_json::Map<String, Value>> = 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<String, Value>| !m.is_empty());
|
||||||
|
|
||||||
if !is_valid_function_name(&function_name) {
|
if !is_valid_function_name(&function_name) {
|
||||||
let error = ErrorData {
|
let error = ErrorData {
|
||||||
code: ErrorCode::INVALID_REQUEST,
|
code: ErrorCode::INVALID_REQUEST,
|
||||||
@@ -346,11 +378,15 @@ pub fn response_to_message(response: &Value) -> anyhow::Result<Message> {
|
|||||||
)),
|
)),
|
||||||
data: None,
|
data: None,
|
||||||
};
|
};
|
||||||
content.push(MessageContent::tool_request(id, Err(error)));
|
content.push(MessageContent::tool_request_with_metadata(
|
||||||
|
id,
|
||||||
|
Err(error),
|
||||||
|
metadata.as_ref(),
|
||||||
|
));
|
||||||
} else {
|
} else {
|
||||||
match safely_parse_json(&arguments_str) {
|
match safely_parse_json(&arguments_str) {
|
||||||
Ok(params) => {
|
Ok(params) => {
|
||||||
content.push(MessageContent::tool_request(
|
content.push(MessageContent::tool_request_with_metadata(
|
||||||
id,
|
id,
|
||||||
Ok(CallToolRequestParams {
|
Ok(CallToolRequestParams {
|
||||||
meta: None,
|
meta: None,
|
||||||
@@ -358,6 +394,7 @@ pub fn response_to_message(response: &Value) -> anyhow::Result<Message> {
|
|||||||
name: function_name.into(),
|
name: function_name.into(),
|
||||||
arguments: Some(object(params)),
|
arguments: Some(object(params)),
|
||||||
}),
|
}),
|
||||||
|
metadata.as_ref(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -369,7 +406,11 @@ pub fn response_to_message(response: &Value) -> anyhow::Result<Message> {
|
|||||||
)),
|
)),
|
||||||
data: None,
|
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() {
|
if chunk.choices.is_empty() {
|
||||||
yield (None, usage)
|
yield (None, usage)
|
||||||
} else if chunk.choices[0].delta.tool_calls.as_ref().is_some_and(|tc| !tc.is_empty()) {
|
} 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<i32, (String, String, String)> = std::collections::HashMap::new();
|
let mut tool_call_data: ToolCallData = HashMap::new();
|
||||||
|
|
||||||
if let Some(tool_calls) = &chunk.choices[0].delta.tool_calls {
|
if let Some(tool_calls) = &chunk.choices[0].delta.tool_calls {
|
||||||
for tool_call in 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) {
|
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 {
|
if let Some(delta_tool_calls) = &tool_chunk.choices[0].delta.tool_calls {
|
||||||
for delta_call in delta_tool_calls {
|
for delta_call in delta_tool_calls {
|
||||||
if let Some(index) = delta_call.index {
|
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);
|
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) {
|
} 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<ProviderMetadata> = if !accumulated_reasoning.is_empty() {
|
let _metadata: Option<ProviderMetadata> = if !accumulated_reasoning.is_empty() {
|
||||||
let mut map = ProviderMetadata::new();
|
let mut map = ProviderMetadata::new();
|
||||||
map.insert("reasoning_details".to_string(), json!(accumulated_reasoning));
|
map.insert("reasoning_details".to_string(), json!(accumulated_reasoning));
|
||||||
Some(map)
|
Some(map)
|
||||||
@@ -577,23 +625,26 @@ where
|
|||||||
sorted_indices.sort();
|
sorted_indices.sort();
|
||||||
|
|
||||||
for index in sorted_indices {
|
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() {
|
let parsed = if arguments.is_empty() {
|
||||||
Ok(json!({}))
|
Ok(json!({}))
|
||||||
} else {
|
} else {
|
||||||
serde_json::from_str::<Value>(arguments)
|
serde_json::from_str::<Value>(arguments)
|
||||||
};
|
};
|
||||||
|
|
||||||
|
let metadata = extra_fields.as_ref().filter(|m| !m.is_empty());
|
||||||
|
|
||||||
let content = match parsed {
|
let content = match parsed {
|
||||||
Ok(params) => {
|
Ok(params) => {
|
||||||
MessageContent::tool_request_with_metadata(
|
MessageContent::tool_request_with_metadata(
|
||||||
id.clone(),
|
id.clone(),
|
||||||
Ok(CallToolRequestParams {
|
Ok(CallToolRequestParams {
|
||||||
meta: None, task: None,
|
meta: None,
|
||||||
|
task: None,
|
||||||
name: function_name.clone().into(),
|
name: function_name.clone().into(),
|
||||||
arguments: Some(object(params))
|
arguments: Some(object(params)),
|
||||||
}),
|
}),
|
||||||
metadata.as_ref(),
|
metadata,
|
||||||
)
|
)
|
||||||
},
|
},
|
||||||
Err(e) => {
|
Err(e) => {
|
||||||
@@ -605,7 +656,7 @@ where
|
|||||||
)),
|
)),
|
||||||
data: None,
|
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);
|
contents.push(content);
|
||||||
@@ -1670,4 +1721,117 @@ data: [DONE]
|
|||||||
|
|
||||||
Ok(())
|
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");
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user