summaryrefslogtreecommitdiffstats
path: root/src/config
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-17 09:14:54 +0800
committerGitHub <noreply@github.com>2024-06-17 09:14:54 +0800
commit638bf3276613265da761b4d74f269fa77c8f4bb1 (patch)
tree9667711b92bc3664ad72c179cfe1cf4bbf54d322 /src/config
parent12872b3d2956a26fb7d1cd464b3672ebe40d94c1 (diff)
downloadaichat-638bf3276613265da761b4d74f269fa77c8f4bb1.tar.gz
refactor: improve code quatity (#604)
Diffstat (limited to 'src/config')
-rw-r--r--src/config/input.rs8
-rw-r--r--src/config/mod.rs126
2 files changed, 73 insertions, 61 deletions
diff --git a/src/config/input.rs b/src/config/input.rs
index 0c93c1f..48b0359 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -4,7 +4,7 @@ use crate::client::{
init_client, ChatCompletionsData, Client, ImageUrl, Message, MessageContent,
MessageContentPart, MessageRole, Model,
};
-use crate::function::{ToolCallResult, ToolResults};
+use crate::function::{ToolResult, ToolResults};
use crate::utils::{base64_encode, sha256, AbortSignal};
use anyhow::{bail, Context, Result};
@@ -154,11 +154,7 @@ impl Input {
self.patched_text.take();
}
- pub fn merge_tool_call(
- mut self,
- output: String,
- tool_call_results: Vec<ToolCallResult>,
- ) -> Self {
+ pub fn merge_tool_call(mut self, output: String, tool_call_results: Vec<ToolResult>) -> Self {
match self.tool_call.as_mut() {
Some(exist_tool_call_results) => {
exist_tool_call_results.0.extend(tool_call_results);
diff --git a/src/config/mod.rs b/src/config/mod.rs
index a905ae8..efbf984 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -12,7 +12,7 @@ use crate::client::{
create_client_config, list_chat_models, list_client_types, ClientConfig, Model,
OPENAI_COMPATIBLE_PLATFORMS,
};
-use crate::function::{FunctionDeclaration, Functions, FunctionsFilter, ToolCallResult};
+use crate::function::{FunctionDeclaration, Functions, FunctionsFilter, ToolResult};
use crate::rag::Rag;
use crate::render::{MarkdownRender, RenderOptions};
use crate::utils::*;
@@ -229,60 +229,6 @@ impl Config {
Ok(path)
}
- pub fn save_message(
- &mut self,
- input: &mut Input,
- output: &str,
- tool_call_results: &[ToolCallResult],
- ) -> Result<()> {
- input.clear_patch_text();
- self.last_message = Some((input.clone(), output.to_string(), self.bot.is_some()));
-
- if self.dry_run || output.is_empty() || !tool_call_results.is_empty() {
- return Ok(());
- }
-
- if let Some(session) = input.session_mut(&mut self.session) {
- session.add_message(input, output)?;
- return Ok(());
- }
-
- if !self.save {
- return Ok(());
- }
- let mut file = self.open_message_file()?;
- if output.is_empty() || !self.save {
- return Ok(());
- }
- let timestamp = now();
- let summary = input.summary();
- let input_markdown = input.render();
- let scope = if self.bot.is_none() {
- let role_name = if input.role().is_derived() {
- None
- } else {
- Some(input.role().name())
- };
- match (role_name, input.rag_name()) {
- (Some(role), Some(rag_name)) => format!(" ({role}#{rag_name})"),
- (Some(role), _) => format!(" ({role})"),
- (None, Some(rag_name)) => format!(" (#{rag_name})"),
- _ => String::new(),
- }
- } else {
- String::new()
- };
- let output = format!("# CHAT: {summary} [{timestamp}]{scope}\n{input_markdown}\n--------\n{output}\n--------\n\n",);
- file.write_all(output.as_bytes())
- .with_context(|| "Failed to save message")
- }
-
- pub fn maybe_copy(&self, text: &str) {
- if self.auto_copy {
- let _ = set_text(text);
- }
- }
-
pub fn config_file() -> Result<PathBuf> {
match env::var(get_env_name("config_file")) {
Ok(value) => Ok(PathBuf::from(value)),
@@ -838,6 +784,7 @@ impl Config {
if let Some(session) = self.session.as_mut() {
session.set_compressing(false);
}
+ self.last_message = None;
}
pub async fn use_rag(
@@ -1246,6 +1193,75 @@ impl Config {
output
}
+ pub fn before_chat_completion(&mut self, input: &Input) -> Result<()> {
+ self.last_message = Some((input.clone(), String::new(), self.bot.is_some()));
+ Ok(())
+ }
+
+ pub fn after_chat_completion(
+ &mut self,
+ input: &mut Input,
+ output: &str,
+ tool_results: &[ToolResult],
+ ) -> Result<()> {
+ input.clear_patch_text();
+ self.last_message = Some((input.clone(), output.to_string(), self.bot.is_some()));
+ self.save_message(input, output, tool_results)?;
+ self.maybe_copy(output);
+ Ok(())
+ }
+
+ fn save_message(
+ &mut self,
+ input: &mut Input,
+ output: &str,
+ tool_results: &[ToolResult],
+ ) -> Result<()> {
+ if self.dry_run || output.is_empty() || !tool_results.is_empty() {
+ return Ok(());
+ }
+
+ if let Some(session) = input.session_mut(&mut self.session) {
+ session.add_message(input, output)?;
+ return Ok(());
+ }
+
+ if !self.save {
+ return Ok(());
+ }
+ let mut file = self.open_message_file()?;
+ if output.is_empty() || !self.save {
+ return Ok(());
+ }
+ let timestamp = now();
+ let summary = input.summary();
+ let input_markdown = input.render();
+ let scope = if self.bot.is_none() {
+ let role_name = if input.role().is_derived() {
+ None
+ } else {
+ Some(input.role().name())
+ };
+ match (role_name, input.rag_name()) {
+ (Some(role), Some(rag_name)) => format!(" ({role}#{rag_name})"),
+ (Some(role), _) => format!(" ({role})"),
+ (None, Some(rag_name)) => format!(" (#{rag_name})"),
+ _ => String::new(),
+ }
+ } else {
+ String::new()
+ };
+ let output = format!("# CHAT: {summary} [{timestamp}]{scope}\n{input_markdown}\n--------\n{output}\n--------\n\n",);
+ file.write_all(output.as_bytes())
+ .with_context(|| "Failed to save message")
+ }
+
+ fn maybe_copy(&self, text: &str) {
+ if self.auto_copy {
+ let _ = set_text(text);
+ }
+ }
+
fn open_message_file(&self) -> Result<File> {
let path = self.messages_file()?;
ensure_parent_exists(&path)?;