summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-11-13 18:24:43 +0800
committerGitHub <noreply@github.com>2024-11-13 18:24:43 +0800
commitff0ea19b48a18e9a1849bd2a653173a3ed4e4560 (patch)
treeb8a8b9d15639649f2fa6e7dc72adcfae770bfc4e /src
parent163ab626cd2b8875161b2598e9400df201eacaa8 (diff)
downloadaichat-ff0ea19b48a18e9a1849bd2a653173a3ed4e4560.tar.gz
fix: invalid request on qianwen multi tool-calls (#993)
Diffstat (limited to 'src')
-rw-r--r--src/client/bedrock.rs5
-rw-r--r--src/client/claude.rs5
-rw-r--r--src/client/cohere.rs4
-rw-r--r--src/client/ernie.rs4
-rw-r--r--src/client/openai.rs9
-rw-r--r--src/client/vertexai.rs3
-rw-r--r--src/config/input.rs5
-rw-r--r--src/function.rs25
-rw-r--r--src/serve.rs2
9 files changed, 47 insertions, 15 deletions
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs
index f7b3019..e1d6c65 100644
--- a/src/client/bedrock.rs
+++ b/src/client/bedrock.rs
@@ -363,7 +363,10 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu
"content": content,
})]
}
- MessageContent::ToolResults((tool_results, text)) => {
+ MessageContent::ToolResults(results) => {
+ let ToolResults {
+ tool_results, text, ..
+ } = results;
let mut assistant_parts = vec![];
let mut user_parts = vec![];
if !text.is_empty() {
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 7b472a7..0c905b1 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -202,7 +202,10 @@ pub fn claude_build_chat_completions_body(
"content": content,
})]
}
- MessageContent::ToolResults((tool_results, text)) => {
+ MessageContent::ToolResults(results) => {
+ let ToolResults {
+ tool_results, text, ..
+ } = results;
let mut assistant_parts = vec![];
let mut user_parts = vec![];
if !text.is_empty() {
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index f471e7a..c5a31f5 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -213,8 +213,8 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu
.collect();
Some(json!({ "role": role, "message": list.join("\n\n") }))
}
- MessageContent::ToolResults((results, _)) => {
- tool_results = Some(results);
+ MessageContent::ToolResults(results) => {
+ tool_results = Some(results.tool_results);
None
}
}
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index d7f1ffb..0c325e5 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -231,9 +231,9 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Valu
.flat_map(|message| {
let Message { role, content } = message;
match content {
- MessageContent::ToolResults((tool_results, _)) => {
+ MessageContent::ToolResults(results) => {
let mut list = vec![];
- for tool_result in tool_results {
+ for tool_result in results.tool_results {
list.push(json!({
"role": "assistant",
"content": format!("Action: {}\nAction Input: {}", tool_result.call.name, tool_result.call.arguments)
diff --git a/src/client/openai.rs b/src/client/openai.rs
index c4c2b0c..0f01e36 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -205,8 +205,13 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod
.flat_map(|message| {
let Message { role, content } = message;
match content {
- MessageContent::ToolResults((tool_results, text)) => {
- if let Some(true) = tool_results.first().map(|v| v.call.id.is_some()) {
+ MessageContent::ToolResults(results) => {
+ let ToolResults {
+ tool_results,
+ text,
+ sequence,
+ } = results;
+ if !sequence {
let tool_calls: Vec<_> = tool_results.iter().map(|tool_result| {
json!({
"id": tool_result.call.id,
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 7ba205e..452eb09 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -342,7 +342,8 @@ pub fn gemini_build_chat_completions_body(
.collect();
vec![json!({ "role": role, "parts": parts })]
},
- MessageContent::ToolResults((tool_results, _)) => {
+ MessageContent::ToolResults(results) => {
+ let tool_results = results.tool_results;
let model_parts: Vec<Value> = tool_results.iter().map(|tool_result| {
json!({
"functionCall": {
diff --git a/src/config/input.rs b/src/config/input.rs
index 54b2bb0..bc58f32 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -186,10 +186,9 @@ impl Input {
pub fn merge_tool_call(mut self, output: String, tool_results: Vec<ToolResult>) -> Self {
match self.tool_call.as_mut() {
Some(exist_tool_results) => {
- exist_tool_results.0.extend(tool_results);
- exist_tool_results.1 = output;
+ exist_tool_results.extend(tool_results, output);
}
- None => self.tool_call = Some((tool_results, output)),
+ None => self.tool_call = Some(ToolResults::new(tool_results, output)),
}
self
}
diff --git a/src/function.rs b/src/function.rs
index 517b9eb..f77467b 100644
--- a/src/function.rs
+++ b/src/function.rs
@@ -13,8 +13,6 @@ use std::{
path::{Path, PathBuf},
};
-pub type ToolResults = (Vec<ToolResult>, String);
-
#[cfg(windows)]
const PATH_SEP: &str = ";";
#[cfg(not(windows))]
@@ -268,6 +266,29 @@ impl ToolCall {
}
}
+#[derive(Debug, Clone, Deserialize, Serialize)]
+pub struct ToolResults {
+ pub tool_results: Vec<ToolResult>,
+ pub text: String,
+ pub sequence: bool,
+}
+
+impl ToolResults {
+ pub fn new(tool_results: Vec<ToolResult>, text: String) -> Self {
+ Self {
+ tool_results,
+ text,
+ sequence: false,
+ }
+ }
+
+ pub fn extend(&mut self, tool_results: Vec<ToolResult>, _text: String) {
+ self.tool_results.extend(tool_results);
+ self.text.clear();
+ self.sequence = true;
+ }
+}
+
#[cfg(windows)]
fn polyfill_cmd_name<T: AsRef<Path>>(cmd_name: &str, bin_dir: &[T]) -> String {
let cmd_name = cmd_name.to_string();
diff --git a/src/serve.rs b/src/serve.rs
index 0b4c475..84e6c87 100644
--- a/src/serve.rs
+++ b/src/serve.rs
@@ -889,7 +889,7 @@ fn parse_messages(message: Vec<Value>) -> Result<Vec<Message>> {
}
output.push(Message::new(
MessageRole::Assistant,
- MessageContent::ToolResults((list, text)),
+ MessageContent::ToolResults(ToolResults::new(list, text)),
));
tool_results = None;
} else {