diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-17 09:14:54 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-17 09:14:54 +0800 |
| commit | 638bf3276613265da761b4d74f269fa77c8f4bb1 (patch) | |
| tree | 9667711b92bc3664ad72c179cfe1cf4bbf54d322 /src/config | |
| parent | 12872b3d2956a26fb7d1cd464b3672ebe40d94c1 (diff) | |
| download | aichat-638bf3276613265da761b4d74f269fa77c8f4bb1.tar.gz | |
refactor: improve code quatity (#604)
Diffstat (limited to 'src/config')
| -rw-r--r-- | src/config/input.rs | 8 | ||||
| -rw-r--r-- | src/config/mod.rs | 126 |
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)?; |
