summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-11-14 08:14:32 +0800
committerGitHub <noreply@github.com>2024-11-14 08:14:32 +0800
commit80684ec84047bda0299f8fd845b1c97f6fac5c99 (patch)
treef37e377accd664d31a2bfe977521cfc2157a9cd1 /src/config
parentcfa9217422dfa6cf10bed6e6e3fab0722f58588d (diff)
downloadaichat-80684ec84047bda0299f8fd845b1c97f6fac5c99.tar.gz
refactor: improve tool calls (#995)
- rename MessageContent:ToolResults to MessageContent:ToolCalls - rename ToolResults to MessageContentToolCalls - persist tool_calls to messages.md
Diffstat (limited to 'src/config')
-rw-r--r--src/config/input.rs26
-rw-r--r--src/config/mod.rs18
-rw-r--r--src/config/session.rs4
3 files changed, 31 insertions, 17 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
index 55a4d45..18b192b 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -2,9 +2,9 @@ use super::*;
use crate::client::{
init_client, patch_system_message, ChatCompletionsData, Client, ImageUrl, Message,
- MessageContent, MessageContentPart, MessageRole, Model,
+ MessageContent, MessageContentPart, MessageContentToolCalls, MessageRole, Model,
};
-use crate::function::{ToolResult, ToolResults};
+use crate::function::ToolResult;
use crate::utils::{base64_encode, sha256, AbortSignal};
use anyhow::{bail, Context, Result};
@@ -29,7 +29,7 @@ pub struct Input {
regenerate: bool,
medias: Vec<String>,
data_urls: HashMap<String, String>,
- tool_results: Option<ToolResults>,
+ tool_calls: Option<MessageContentToolCalls>,
rag_name: Option<String>,
role: Role,
with_session: bool,
@@ -48,7 +48,7 @@ impl Input {
regenerate: false,
medias: Default::default(),
data_urls: Default::default(),
- tool_results: None,
+ tool_calls: None,
rag_name: None,
role,
with_session,
@@ -104,7 +104,7 @@ impl Input {
regenerate: false,
medias,
data_urls,
- tool_results: Default::default(),
+ tool_calls: Default::default(),
rag_name: None,
role,
with_session,
@@ -120,8 +120,8 @@ impl Input {
self.data_urls.clone()
}
- pub fn tool_results(&self) -> &Option<ToolResults> {
- &self.tool_results
+ pub fn tool_calls(&self) -> &Option<MessageContentToolCalls> {
+ &self.tool_calls
}
pub fn text(&self) -> String {
@@ -187,12 +187,12 @@ impl Input {
self.rag_name.as_deref()
}
- pub fn merge_tool_call(mut self, output: String, tool_results: Vec<ToolResult>) -> Self {
- match self.tool_results.as_mut() {
+ pub fn merge_tool_results(mut self, output: String, tool_results: Vec<ToolResult>) -> Self {
+ match self.tool_calls.as_mut() {
Some(exist_tool_results) => {
- exist_tool_results.extend(tool_results, output);
+ exist_tool_results.merge(tool_results, output);
}
- None => self.tool_results = Some(ToolResults::new(tool_results, output)),
+ None => self.tool_calls = Some(MessageContentToolCalls::new(tool_results, output)),
}
self
}
@@ -232,10 +232,10 @@ impl Input {
} else {
self.role().build_messages(self)
};
- if let Some(tool_results) = &self.tool_results {
+ if let Some(tool_calls) = &self.tool_calls {
messages.push(Message::new(
MessageRole::Assistant,
- MessageContent::ToolResults(tool_results.clone()),
+ MessageContent::ToolCalls(tool_calls.clone()),
))
}
Ok(messages)
diff --git a/src/config/mod.rs b/src/config/mod.rs
index 9da8dae..136fee1 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -10,7 +10,7 @@ use self::session::Session;
use crate::client::{
create_client_config, list_chat_models, list_client_types, list_reranker_models, ClientConfig,
- Model, OPENAI_COMPATIBLE_PLATFORMS,
+ MessageContentToolCalls, Model, OPENAI_COMPATIBLE_PLATFORMS,
};
use crate::function::{FunctionDeclaration, Functions, ToolResult};
use crate::rag::Rag;
@@ -1863,8 +1863,22 @@ impl Config {
} else {
String::new()
};
+ let tool_calls = match input.tool_calls() {
+ Some(MessageContentToolCalls {
+ tool_results, text, ..
+ }) => {
+ let mut lines = vec!["<tool_calls>".to_string()];
+ if !text.is_empty() {
+ lines.push(text.clone());
+ }
+ lines.push(serde_json::to_string(&tool_results).unwrap_or_default());
+ lines.push("</tool_calls>\n".to_string());
+ lines.join("\n")
+ }
+ None => String::new(),
+ };
let output = format!(
- "# CHAT: {summary} [{timestamp}]{scope}\n{raw_input}\n--------\n{output}\n--------\n\n",
+ "# CHAT: {summary} [{timestamp}]{scope}\n{raw_input}\n--------\n{tool_calls}{output}\n--------\n\n",
);
file.write_all(output.as_bytes())
.with_context(|| "Failed to save message")
diff --git a/src/config/session.rs b/src/config/session.rs
index 34d6393..633f621 100644
--- a/src/config/session.rs
+++ b/src/config/session.rs
@@ -419,10 +419,10 @@ impl Session {
.push(Message::new(MessageRole::User, input.message_content()));
}
self.data_urls.extend(input.data_urls());
- if let Some(tool_results) = input.tool_results() {
+ if let Some(tool_calls) = input.tool_calls() {
self.messages.push(Message::new(
MessageRole::Tool,
- MessageContent::ToolResults(tool_results.clone()),
+ MessageContent::ToolCalls(tool_calls.clone()),
))
}
self.messages.push(Message::new(