summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/bedrock.rs7
-rw-r--r--src/client/claude.rs7
-rw-r--r--src/client/cohere.rs4
-rw-r--r--src/client/ernie.rs6
-rw-r--r--src/client/message.rs40
-rw-r--r--src/client/mod.rs2
-rw-r--r--src/client/model.rs9
-rw-r--r--src/client/openai.rs5
-rw-r--r--src/client/vertexai.rs3
9 files changed, 50 insertions, 33 deletions
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs
index e1d6c65..78f9d2b 100644
--- a/src/client/bedrock.rs
+++ b/src/client/bedrock.rs
@@ -363,10 +363,9 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Resu
"content": content,
})]
}
- MessageContent::ToolResults(results) => {
- let ToolResults {
- tool_results, text, ..
- } = results;
+ MessageContent::ToolCalls(MessageContentToolCalls {
+ tool_results, text, ..
+ }) => {
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 0c905b1..c7ef88a 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -202,10 +202,9 @@ pub fn claude_build_chat_completions_body(
"content": content,
})]
}
- MessageContent::ToolResults(results) => {
- let ToolResults {
- tool_results, text, ..
- } = results;
+ MessageContent::ToolCalls(MessageContentToolCalls {
+ tool_results, text, ..
+ }) => {
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 c5a31f5..9726300 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.tool_results);
+ MessageContent::ToolCalls(tool_calls) => {
+ tool_results = Some(tool_calls.tool_results);
None
}
}
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 0c325e5..7e1f8e7 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -231,9 +231,11 @@ fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Valu
.flat_map(|message| {
let Message { role, content } = message;
match content {
- MessageContent::ToolResults(results) => {
+ MessageContent::ToolCalls(MessageContentToolCalls {
+ tool_results, ..
+ }) => {
let mut list = vec![];
- for tool_result in results.tool_results {
+ for tool_result in 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/message.rs b/src/client/message.rs
index 061c0e2..f6b7f6d 100644
--- a/src/client/message.rs
+++ b/src/client/message.rs
@@ -1,6 +1,4 @@
-use super::ToolResults;
-
-use crate::utils::dimmed_text;
+use crate::{function::ToolResult, utils::dimmed_text};
use serde::{Deserialize, Serialize};
@@ -75,7 +73,7 @@ pub enum MessageContent {
Text(String),
Array(Vec<MessageContentPart>),
// Note: This type is primarily for convenience and does not exist in OpenAI's API.
- ToolResults(ToolResults),
+ ToolCalls(MessageContentToolCalls),
}
impl MessageContent {
@@ -103,10 +101,9 @@ impl MessageContent {
}
format!(".file {}{}", files.join(" "), concated_text)
}
- MessageContent::ToolResults(results) => {
- let ToolResults {
- tool_results, text, ..
- } = results;
+ MessageContent::ToolCalls(MessageContentToolCalls {
+ tool_results, text, ..
+ }) => {
let mut lines = vec![];
if !text.is_empty() {
lines.push(text.clone())
@@ -139,7 +136,7 @@ impl MessageContent {
*text = replace_fn(text)
}
}
- MessageContent::ToolResults(_) => {}
+ MessageContent::ToolCalls(_) => {}
}
}
@@ -155,7 +152,7 @@ impl MessageContent {
}
parts.join("\n\n")
}
- MessageContent::ToolResults(_) => String::new(),
+ MessageContent::ToolCalls(_) => String::new(),
}
}
}
@@ -172,6 +169,29 @@ pub struct ImageUrl {
pub url: String,
}
+#[derive(Debug, Clone, Deserialize, Serialize)]
+pub struct MessageContentToolCalls {
+ pub tool_results: Vec<ToolResult>,
+ pub text: String,
+ pub sequence: bool,
+}
+
+impl MessageContentToolCalls {
+ pub fn new(tool_results: Vec<ToolResult>, text: String) -> Self {
+ Self {
+ tool_results,
+ text,
+ sequence: false,
+ }
+ }
+
+ pub fn merge(&mut self, tool_results: Vec<ToolResult>, _text: String) {
+ self.tool_results.extend(tool_results);
+ self.text.clear();
+ self.sequence = true;
+ }
+}
+
pub fn patch_system_message(messages: &mut Vec<Message>) {
if messages[0].role.is_system() {
let system_message = messages.remove(0);
diff --git a/src/client/mod.rs b/src/client/mod.rs
index b22508b..5189f9f 100644
--- a/src/client/mod.rs
+++ b/src/client/mod.rs
@@ -6,7 +6,7 @@ mod macros;
mod model;
mod stream;
-pub use crate::function::{ToolCall, ToolResults};
+pub use crate::function::ToolCall;
pub use crate::utils::PromptKind;
pub use common::*;
pub use message::*;
diff --git a/src/client/model.rs b/src/client/model.rs
index af864e2..5d496d2 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -1,7 +1,7 @@
use super::{
list_chat_models, list_embedding_models, list_reranker_models,
message::{Message, MessageContent, MessageContentPart},
- ToolResults,
+ MessageContentToolCalls,
};
use crate::config::Config;
@@ -237,10 +237,9 @@ impl Model {
MessageContentPart::ImageUrl { .. } => 0,
})
.sum(),
- MessageContent::ToolResults(results) => {
- let ToolResults {
- tool_results, text, ..
- } = results;
+ MessageContent::ToolCalls(MessageContentToolCalls {
+ tool_results, text, ..
+ }) => {
estimate_token_length(text)
+ tool_results
.iter()
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 0f01e36..7f449fb 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -205,12 +205,11 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod
.flat_map(|message| {
let Message { role, content } = message;
match content {
- MessageContent::ToolResults(results) => {
- let ToolResults {
+ MessageContent::ToolCalls(MessageContentToolCalls {
tool_results,
text,
sequence,
- } = results;
+ }) => {
if !sequence {
let tool_calls: Vec<_> = tool_results.iter().map(|tool_result| {
json!({
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 452eb09..07ff2a4 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -342,8 +342,7 @@ pub fn gemini_build_chat_completions_body(
.collect();
vec![json!({ "role": role, "parts": parts })]
},
- MessageContent::ToolResults(results) => {
- let tool_results = results.tool_results;
+ MessageContent::ToolCalls(MessageContentToolCalls { tool_results, .. }) => {
let model_parts: Vec<Value> = tool_results.iter().map(|tool_result| {
json!({
"functionCall": {